diff --git a/.github/workflows/Release.yml b/.github/workflows/Release.yml index 2791b19e2..7933f62af 100644 --- a/.github/workflows/Release.yml +++ b/.github/workflows/Release.yml @@ -9,7 +9,18 @@ on: permissions: write-all # Necessary for the generate-build-provenance action with containers jobs: + stress-tests: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: stable + - name: Run release stress tests + run: go test -tags=stress -count=1 ./cmd/cloud-init-server ./pkg/wgtunnel ./internal/memstore ./internal/smdclient + release: + needs: stress-tests uses: OpenCHAMI/github-actions/.github/workflows/go-build-release.yml@v3.2 with: cgo-enabled: "1" diff --git a/cmd/cloud-init-server/handlers.go b/cmd/cloud-init-server/handlers.go index 3b7984176..bebf91264 100644 --- a/cmd/cloud-init-server/handlers.go +++ b/cmd/cloud-init-server/handlers.go @@ -5,12 +5,11 @@ import ( "io" "net/http" - // Import to run swag.Register() to generated docs "github.com/go-chi/chi/v5" + // Import to run swag.Register() to generated docs _ "github.com/openchami/cloud-init/docs" "github.com/openchami/cloud-init/internal/smdclient" "github.com/openchami/cloud-init/pkg/cistore" - "github.com/openchami/cloud-init/pkg/wgtunnel" "github.com/rs/zerolog/log" "github.com/swaggo/swag" ) @@ -207,7 +206,7 @@ func InstanceInfoHandler(sm smdclient.SMDClientInterface, store cistore.Store) h // @Param hostname formData string true "Node's given hostname" // @Param fqdn formData string true "Node's given fully-qualified domain name" // @Router /phone-home/{id} [post] -func PhoneHomeHandler(wg *wgtunnel.InterfaceManager, sm smdclient.SMDClientInterface) http.HandlerFunc { +func PhoneHomeHandler(peerRemovalQueue *PeerRemovalQueue, sm smdclient.SMDClientInterface) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { w.WriteHeader(http.StatusMethodNotAllowed) @@ -249,12 +248,10 @@ func PhoneHomeHandler(wg *wgtunnel.InterfaceManager, sm smdclient.SMDClientInter Msgf("Received phone home data: pub_key_rsa=%s, pub_key_ecdsa=%s, pub_key_ed25519=%s, instance_id=%s, hostname=%s, fqdn=%s", pubKeyRsa, pubKeyEcdsa, pubKeyEd25519, instanceId, hostname, fqdn) - if wg != nil { - go func() { - _ = wg.RemovePeer(peerName) // Explicitly ignoring the error here. There's nothing to do with it within the goroutine. - }() - - w.WriteHeader(http.StatusOK) + if peerRemovalQueue != nil && !peerRemovalQueue.TryEnqueue(peerName) { + http.Error(w, "WireGuard peer removal queue is full", http.StatusServiceUnavailable) + return } + w.WriteHeader(http.StatusOK) } } diff --git a/cmd/cloud-init-server/main.go b/cmd/cloud-init-server/main.go index 7425ee01c..5f30e499e 100644 --- a/cmd/cloud-init-server/main.go +++ b/cmd/cloud-init-server/main.go @@ -269,6 +269,10 @@ func startServer() error { // Create router router := chi.NewRouter() + var peerRemovalQueue *PeerRemovalQueue + if wgInterfaceManager != nil { + peerRemovalQueue = NewPeerRemovalQueue(wgInterfaceManager) + } // Add middleware router.Use( @@ -282,7 +286,7 @@ func startServer() error { ) // Setup routes - initCiClientRouter(router, handler, wgInterfaceManager) + initCiClientRouter(router, handler, wgInterfaceManager, peerRemovalQueue) initCiAdminRouter(router, handler) // Add secure routes if JWKS is configured @@ -315,7 +319,7 @@ func parseBool(str string) bool { return strings.EqualFold(str, "true") || str == "1" } -func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManager *wgtunnel.InterfaceManager) { +func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManager *wgtunnel.InterfaceManager, peerRemovalQueue *PeerRemovalQueue) { // Add cloud-init endpoints to router router.Get("/openapi.json", DocsHandler) router.Get("/version", VersionHandler) @@ -330,7 +334,7 @@ func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManage router.Get("/vendor-data", VendorDataHandler(handler.sm, handler.store, baseUrl)) router.Get("/{group}.yaml", GroupUserDataHandler(handler.sm, handler.store)) } - router.Post("/phone-home/{id}", PhoneHomeHandler(wgInterfaceManager, handler.sm)) + router.Post("/phone-home/{id}", PhoneHomeHandler(peerRemovalQueue, handler.sm)) router.Post("/wg-init", wgtunnel.AddClientHandler(wgInterfaceManager, handler.sm)) } diff --git a/cmd/cloud-init-server/peer_removal_queue.go b/cmd/cloud-init-server/peer_removal_queue.go new file mode 100644 index 000000000..077b9f18c --- /dev/null +++ b/cmd/cloud-init-server/peer_removal_queue.go @@ -0,0 +1,52 @@ +package main + +import "github.com/rs/zerolog/log" + +const ( + defaultPeerRemovalWorkers = 2 + defaultPeerRemovalBuffer = 64 +) + +type peerRemover interface { + RemovePeer(peerName string) error +} + +type PeerRemovalQueue struct { + remover peerRemover + jobs chan string +} + +func NewPeerRemovalQueue(remover peerRemover) *PeerRemovalQueue { + return newPeerRemovalQueue(remover, defaultPeerRemovalWorkers, defaultPeerRemovalBuffer) +} + +func newPeerRemovalQueue(remover peerRemover, workers int, buffer int) *PeerRemovalQueue { + queue := &PeerRemovalQueue{ + remover: remover, + jobs: make(chan string, buffer), + } + for range workers { + go queue.work() + } + return queue +} + +func (q *PeerRemovalQueue) TryEnqueue(peerName string) bool { + if q == nil || q.remover == nil { + return true + } + select { + case q.jobs <- peerName: + return true + default: + return false + } +} + +func (q *PeerRemovalQueue) work() { + for peerName := range q.jobs { + if err := q.remover.RemovePeer(peerName); err != nil { + log.Error().Err(err).Str("peer", peerName).Msg("failed to remove WireGuard peer") + } + } +} diff --git a/cmd/cloud-init-server/peer_removal_queue_stress_test.go b/cmd/cloud-init-server/peer_removal_queue_stress_test.go new file mode 100644 index 000000000..f1c0e4c35 --- /dev/null +++ b/cmd/cloud-init-server/peer_removal_queue_stress_test.go @@ -0,0 +1,74 @@ +//go:build stress + +package main + +import ( + "net/http" + "net/http/httptest" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestStressPhoneHomeQueueBackpressure10K(t *testing.T) { + remover := newBlockingPeerRemover() + queue := newPeerRemovalQueue(remover, defaultPeerRemovalWorkers, defaultPeerRemovalBuffer) + handler := PhoneHomeHandler(queue, &phoneHomeSMDClient{}) + + for range defaultPeerRemovalWorkers { + recorder := httptest.NewRecorder() + handler(recorder, phoneHomeRequest(t)) + if recorder.Code != http.StatusOK { + t.Fatalf("worker-fill response status = %d, want %d", recorder.Code, http.StatusOK) + } + } + for range defaultPeerRemovalWorkers { + select { + case <-remover.started: + case <-time.After(time.Second): + t.Fatal("worker did not start removal") + } + } + + const requestCount = 10_000 + var okCount atomic.Int64 + var unavailableCount atomic.Int64 + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(requestCount) + start.Add(1) + done.Add(requestCount) + + for range requestCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + + recorder := httptest.NewRecorder() + handler(recorder, phoneHomeRequest(t)) + switch recorder.Code { + case http.StatusOK: + okCount.Add(1) + case http.StatusServiceUnavailable: + unavailableCount.Add(1) + default: + t.Errorf("response status = %d, want %d or %d", recorder.Code, http.StatusOK, http.StatusServiceUnavailable) + } + }() + } + + ready.Wait() + start.Done() + done.Wait() + + if got := okCount.Load(); got != defaultPeerRemovalBuffer { + t.Fatalf("accepted removals = %d, want %d", got, defaultPeerRemovalBuffer) + } + if got := unavailableCount.Load(); got != requestCount-defaultPeerRemovalBuffer { + t.Fatalf("backpressured removals = %d, want %d", got, requestCount-defaultPeerRemovalBuffer) + } + close(remover.release) +} diff --git a/cmd/cloud-init-server/peer_removal_queue_test.go b/cmd/cloud-init-server/peer_removal_queue_test.go new file mode 100644 index 000000000..914901916 --- /dev/null +++ b/cmd/cloud-init-server/peer_removal_queue_test.go @@ -0,0 +1,126 @@ +package main + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/go-chi/chi/v5" + "github.com/openchami/cloud-init/internal/smdclient" +) + +type blockingPeerRemover struct { + started chan struct{} + release chan struct{} + removed chan string +} + +func newBlockingPeerRemover() *blockingPeerRemover { + return &blockingPeerRemover{ + started: make(chan struct{}, 16), + release: make(chan struct{}), + removed: make(chan string, 16), + } +} + +func (r *blockingPeerRemover) RemovePeer(peerName string) error { + r.started <- struct{}{} + <-r.release + r.removed <- peerName + return nil +} + +type phoneHomeSMDClient struct { + smdclient.FakeSMDClient +} + +func (phoneHomeSMDClient) IDfromIP(string) (string, error) { + return "x0c0s0b0n0", nil +} + +func (phoneHomeSMDClient) IPfromID(string) (string, error) { + return "10.1.0.1", nil +} + +func TestPeerRemovalQueueBoundsWork(t *testing.T) { + remover := newBlockingPeerRemover() + queue := newPeerRemovalQueue(remover, 1, 1) + + if !queue.TryEnqueue("peer-1") { + t.Fatal("first enqueue unexpectedly failed") + } + select { + case <-remover.started: + case <-time.After(time.Second): + t.Fatal("worker did not start first removal") + } + if !queue.TryEnqueue("peer-2") { + t.Fatal("buffered enqueue unexpectedly failed") + } + if queue.TryEnqueue("peer-3") { + t.Fatal("enqueue succeeded when worker and buffer were saturated") + } + + close(remover.release) + for range 2 { + select { + case <-remover.removed: + case <-time.After(time.Second): + t.Fatal("queued removal did not finish") + } + } +} + +func TestPhoneHomeHandlerReturnsUnavailableWhenRemovalQueueFull(t *testing.T) { + remover := newBlockingPeerRemover() + queue := newPeerRemovalQueue(remover, 1, 1) + handler := PhoneHomeHandler(queue, &phoneHomeSMDClient{}) + + first := httptest.NewRecorder() + handler(first, phoneHomeRequest(t)) + if first.Code != http.StatusOK { + t.Fatalf("first response status = %d, want %d", first.Code, http.StatusOK) + } + select { + case <-remover.started: + case <-time.After(time.Second): + t.Fatal("worker did not start first removal") + } + + second := httptest.NewRecorder() + handler(second, phoneHomeRequest(t)) + if second.Code != http.StatusOK { + t.Fatalf("second response status = %d, want %d", second.Code, http.StatusOK) + } + + third := httptest.NewRecorder() + handler(third, phoneHomeRequest(t)) + if third.Code != http.StatusServiceUnavailable { + t.Fatalf("third response status = %d, want %d", third.Code, http.StatusServiceUnavailable) + } + close(remover.release) +} + +func TestPhoneHomeHandlerWithoutWireGuardStillReturnsOK(t *testing.T) { + handler := PhoneHomeHandler(nil, &phoneHomeSMDClient{}) + recorder := httptest.NewRecorder() + handler(recorder, phoneHomeRequest(t)) + if recorder.Code != http.StatusOK { + t.Fatalf("response status = %d, want %d", recorder.Code, http.StatusOK) + } +} + +func phoneHomeRequest(t *testing.T) *http.Request { + t.Helper() + r := httptest.NewRequest(http.MethodPost, "/phone-home/x0c0s0b0n0", nil) + r.RemoteAddr = "10.1.0.1:12345" + rctx := chi.NewRouteContext() + rctx.URLParams.Add("id", "x0c0s0b0n0") + return r.WithContext(contextWithRoute(r.Context(), rctx)) +} + +func contextWithRoute(ctx context.Context, rctx *chi.Context) context.Context { + return context.WithValue(ctx, chi.RouteCtxKey, rctx) +} diff --git a/internal/memstore/ciMemStore.go b/internal/memstore/ciMemStore.go index 03773ec4b..054591ca6 100644 --- a/internal/memstore/ciMemStore.go +++ b/internal/memstore/ciMemStore.go @@ -3,8 +3,10 @@ package memstore import ( "crypto/rand" "fmt" + "maps" "os" "path/filepath" + "slices" "strings" "sync" @@ -80,7 +82,11 @@ func NewMemStoreFromPath(path string) (*MemStore, error) { func (m *MemStore) GetGroups() map[string]cistore.GroupData { m.GroupsMutex.RLock() defer m.GroupsMutex.RUnlock() - return m.Groups + groups := make(map[string]cistore.GroupData, len(m.Groups)) + for groupName, groupData := range m.Groups { + groups[groupName] = cloneGroupData(groupData) + } + return groups } func (m *MemStore) AddGroupData(groupName string, newGroupData cistore.GroupData) error { @@ -94,7 +100,7 @@ func (m *MemStore) AddGroupData(groupName string, newGroupData cistore.GroupData return fmt.Errorf("group '%s' not added as it already exists", groupName) } else { // does not exist, so create and update - m.Groups[groupName] = newGroupData + m.Groups[groupName] = cloneGroupData(newGroupData) } return nil @@ -106,7 +112,7 @@ func (m *MemStore) GetGroupData(groupName string) (cistore.GroupData, error) { defer m.GroupsMutex.RUnlock() group, ok := m.Groups[groupName] if ok { - return group, nil + return cloneGroupData(group), nil } else { return cistore.GroupData{}, fmt.Errorf("group (%s) not found in memstore", groupName) } @@ -118,13 +124,13 @@ func (m *MemStore) UpdateGroupData(groupName string, groupData cistore.GroupData m.GroupsMutex.Lock() defer m.GroupsMutex.Unlock() if create { - m.Groups[groupName] = groupData + m.Groups[groupName] = cloneGroupData(groupData) return nil } _, ok := m.Groups[groupName] if ok { - m.Groups[groupName] = groupData + m.Groups[groupName] = cloneGroupData(groupData) } else { return fmt.Errorf("group (%s) not found", groupName) } @@ -146,7 +152,7 @@ func (m *MemStore) GetInstanceInfo(nodeName string) (cistore.OpenCHAMIInstanceIn InstanceID: generateInstanceId(), } } - return m.Instances[nodeName], nil + return cloneInstanceInfo(m.Instances[nodeName]), nil } func (m *MemStore) SetInstanceInfo(nodeName string, instanceInfo cistore.OpenCHAMIInstanceInfo) error { @@ -157,11 +163,11 @@ func (m *MemStore) SetInstanceInfo(nodeName string, instanceInfo cistore.OpenCHA if instanceInfo.InstanceID == "" { instanceInfo.InstanceID = generateInstanceId() } - m.Instances[nodeName] = instanceInfo + m.Instances[nodeName] = cloneInstanceInfo(instanceInfo) } else { // This is an update operation. We need to keep the instance ID the same. instanceInfo.InstanceID = m.Instances[nodeName].InstanceID - m.Instances[nodeName] = instanceInfo + m.Instances[nodeName] = cloneInstanceInfo(instanceInfo) } return nil } @@ -176,7 +182,7 @@ func (m *MemStore) DeleteInstanceInfo(nodeName string) error { func (m *MemStore) GetClusterDefaults() (cistore.ClusterDefaults, error) { m.ClusterDefaultsMutex.RLock() defer m.ClusterDefaultsMutex.RUnlock() - return m.ClusterDefaults, nil + return cloneClusterDefaults(m.ClusterDefaults), nil } func (m *MemStore) SetClusterDefaults(clusterDefaults cistore.ClusterDefaults) error { @@ -215,10 +221,53 @@ func (m *MemStore) SetClusterDefaults(clusterDefaults cistore.ClusterDefaults) e log.Debug().Msgf("Setting Public Keys to %v", clusterDefaults.PublicKeys) cd.PublicKeys = clusterDefaults.PublicKeys } - m.ClusterDefaults = cd + m.ClusterDefaults = cloneClusterDefaults(cd) return nil } +func cloneGroupData(groupData cistore.GroupData) cistore.GroupData { + groupData.Data = cloneAnyMap(groupData.Data) + groupData.File.Content = slices.Clone(groupData.File.Content) + groupData.Versions = maps.Clone(groupData.Versions) + return groupData +} + +func cloneInstanceInfo(instanceInfo cistore.OpenCHAMIInstanceInfo) cistore.OpenCHAMIInstanceInfo { + instanceInfo.PublicKeys = slices.Clone(instanceInfo.PublicKeys) + return instanceInfo +} + +func cloneClusterDefaults(clusterDefaults cistore.ClusterDefaults) cistore.ClusterDefaults { + clusterDefaults.PublicKeys = slices.Clone(clusterDefaults.PublicKeys) + return clusterDefaults +} + +func cloneAnyMap(source map[string]any) map[string]any { + if source == nil { + return nil + } + clone := make(map[string]any, len(source)) + for key, value := range source { + clone[key] = cloneAny(value) + } + return clone +} + +func cloneAny(value any) any { + switch typed := value.(type) { + case map[string]any: + return cloneAnyMap(typed) + case []any: + clone := make([]any, len(typed)) + for i, item := range typed { + clone[i] = cloneAny(item) + } + return clone + default: + return value + } +} + func generateInstanceId() string { // in the future, we might want to map the instance-id to an xname or something else. return generateUniqueID("i") diff --git a/internal/memstore/ciMemStore_stress_test.go b/internal/memstore/ciMemStore_stress_test.go new file mode 100644 index 000000000..1c34682b2 --- /dev/null +++ b/internal/memstore/ciMemStore_stress_test.go @@ -0,0 +1,68 @@ +//go:build stress + +package memstore + +import ( + "fmt" + "sync" + "testing" + + "github.com/openchami/cloud-init/pkg/cistore" + "github.com/stretchr/testify/require" +) + +func TestStressMemStoreDefensiveCopies10K(t *testing.T) { + store := NewMemStore() + seed := cistore.GroupData{ + Name: "compute", + Data: map[string]any{ + "nested": map[string]any{"key": "value"}, + "list": []any{"a", map[string]any{"b": "c"}}, + }, + File: cistore.CloudConfigFile{Content: []byte("#cloud-config")}, + Versions: map[string]string{"v1": "one"}, + } + require.NoError(t, store.AddGroupData(seed.Name, seed)) + + const operationCount = 10_000 + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(operationCount) + start.Add(1) + done.Add(operationCount) + + for i := range operationCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + + if i%10 == 0 { + name := fmt.Sprintf("group-%d", i) + _ = store.UpdateGroupData(name, cistore.GroupData{Name: name, Data: map[string]any{"index": i}}, true) + return + } + + groups := store.GetGroups() + if group, found := groups["compute"]; found { + mutateGroupData(group) + } + group, err := store.GetGroupData("compute") + if err == nil { + mutateGroupData(group) + } + }() + } + + ready.Wait() + start.Done() + done.Wait() + + fresh, err := store.GetGroupData("compute") + require.NoError(t, err) + require.Equal(t, "value", fresh.Data["nested"].(map[string]any)["key"]) + require.Equal(t, "c", fresh.Data["list"].([]any)[1].(map[string]any)["b"]) + require.Equal(t, []byte("#cloud-config"), fresh.File.Content) + require.Equal(t, "one", fresh.Versions["v1"]) +} diff --git a/internal/memstore/ciMemStore_test.go b/internal/memstore/ciMemStore_test.go index 81db0a7df..42803fb81 100644 --- a/internal/memstore/ciMemStore_test.go +++ b/internal/memstore/ciMemStore_test.go @@ -158,3 +158,83 @@ func TestConcurrentInstanceAccess(t *testing.T) { wg.Wait() }) } + +func TestMemStoreGroupDataCopies(t *testing.T) { + store := NewMemStore() + original := cistore.GroupData{ + Name: "compute", + Description: "Compute nodes", + Data: map[string]interface{}{ + "role": "compute", + "nested": map[string]interface{}{ + "key": "value", + }, + "list": []interface{}{"a", map[string]interface{}{"b": "c"}}, + }, + File: cistore.CloudConfigFile{ + Content: []byte("#cloud-config"), + Encoding: "plain", + }, + Versions: map[string]string{"v1": "one"}, + } + require.NoError(t, store.AddGroupData(original.Name, original)) + + original.Data["role"] = "mutated" + original.Data["nested"].(map[string]interface{})["key"] = "mutated" + original.Data["list"].([]interface{})[1].(map[string]interface{})["b"] = "mutated" + original.File.Content[0] = '!' + original.Versions["v1"] = "mutated" + + group, err := store.GetGroupData("compute") + require.NoError(t, err) + mutateGroupData(group) + + groups := store.GetGroups() + delete(groups, "compute") + groups = store.GetGroups() + mutateGroupData(groups["compute"]) + + fresh, err := store.GetGroupData("compute") + require.NoError(t, err) + require.Equal(t, "compute", fresh.Data["role"]) + require.Equal(t, "value", fresh.Data["nested"].(map[string]interface{})["key"]) + require.Equal(t, "c", fresh.Data["list"].([]interface{})[1].(map[string]interface{})["b"]) + require.Equal(t, []byte("#cloud-config"), fresh.File.Content) + require.Equal(t, "one", fresh.Versions["v1"]) +} + +func TestMemStoreInstanceAndDefaultsCopies(t *testing.T) { + store := NewMemStore() + instance := cistore.OpenCHAMIInstanceInfo{ + InstanceID: "i-1", + PublicKeys: []string{ + "key-1", + }, + } + require.NoError(t, store.SetInstanceInfo("node1", instance)) + instance.PublicKeys[0] = "mutated" + gotInstance, err := store.GetInstanceInfo("node1") + require.NoError(t, err) + gotInstance.PublicKeys[0] = "mutated-again" + freshInstance, err := store.GetInstanceInfo("node1") + require.NoError(t, err) + require.Equal(t, []string{"key-1"}, freshInstance.PublicKeys) + + defaults := cistore.ClusterDefaults{PublicKeys: []string{"default-key"}} + require.NoError(t, store.SetClusterDefaults(defaults)) + defaults.PublicKeys[0] = "mutated" + gotDefaults, err := store.GetClusterDefaults() + require.NoError(t, err) + gotDefaults.PublicKeys[0] = "mutated-again" + freshDefaults, err := store.GetClusterDefaults() + require.NoError(t, err) + require.Equal(t, []string{"default-key"}, freshDefaults.PublicKeys) +} + +func mutateGroupData(group cistore.GroupData) { + group.Data["role"] = "mutated" + group.Data["nested"].(map[string]interface{})["key"] = "mutated" + group.Data["list"].([]interface{})[1].(map[string]interface{})["b"] = "mutated" + group.File.Content[0] = '!' + group.Versions["v1"] = "mutated" +} diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go index 5f49cc350..c2c62fdf4 100644 --- a/internal/smdclient/SMDclient.go +++ b/internal/smdclient/SMDclient.go @@ -46,7 +46,9 @@ type SMDClient struct { smdBaseURL string tokenEndpoint string accessToken string + accessTokenMutex sync.Mutex nodes map[string]NodeMapping + components map[string]base.Component nodesMutex *sync.RWMutex nodes_last_update time.Time stopCacheRefresh chan struct{} @@ -114,6 +116,7 @@ func NewSMDClient(clusterName, baseurl, jwtURL, accessToken, certPath string, in nodesMutex: &sync.RWMutex{}, nodes_last_update: time.Now(), nodes: make(map[string]NodeMapping), + components: make(map[string]base.Component), stopCacheRefresh: make(chan struct{}), ipToXname: make(map[string]string), macToXname: make(map[string]string), @@ -166,7 +169,7 @@ func (s *SMDClient) ClusterName() string { } // getSMD is a helper function to initialize the SMDClient -func (s *SMDClient) getSMD(ep string, smd interface{}) error { +func (s *SMDClient) getSMD(ep string, smd any) error { url := s.smdBaseURL + ep var resp *http.Response // Manage fetching a new JWT if we initially fail @@ -176,19 +179,21 @@ func (s *SMDClient) getSMD(ep string, smd interface{}) error { if err != nil { return err } - req.Header.Set("Authorization", "Bearer "+s.accessToken) + usedToken := s.currentAccessToken() + req.Header.Set("Authorization", "Bearer "+usedToken) resp, err = s.smdClient.Do(req) if err != nil { return err } if resp.StatusCode == http.StatusUnauthorized { + _ = resp.Body.Close() // Request failed; handle appropriately (based on whether or not // this was a fresh JWT) log.Info().Msg("Cached JWT was rejected by SMD") if !freshToken { log.Info().Msg("Fetching new JWT and retrying...") // Try to refresh the token and retry once - if err2 := s.RefreshToken(); err2 != nil { + if err2 := s.refreshTokenIfCurrent(usedToken); err2 != nil { // If token refresh fails, refresh will attempt again. // While effectively we could ignore the error, it helps // to see why the failure is occurring in case the error @@ -227,6 +232,12 @@ func (s *SMDClient) getSMD(ep string, smd interface{}) error { return nil } +func (s *SMDClient) currentAccessToken() string { + s.accessTokenMutex.Lock() + defer s.accessTokenMutex.Unlock() + return s.accessToken +} + // 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() { @@ -283,6 +294,29 @@ func (s *SMDClient) PopulateNodes() { } } + var componentArray base.ComponentArray + if err := s.getSMD("/hsm/v2/State/Components", &componentArray); err != nil { + log.Error().Err(err).Msg("Failed to get SMD component data") + return + } + nextComponents := make(map[string]base.Component, len(componentArray.Components)) + for _, component := range componentArray.Components { + if component == nil || component.ID == "" { + continue + } + nextComponents[component.ID] = cloneComponent(*component) + } + if len(nextComponents) == 0 { + log.Error().Msg("SMD component data was empty") + return + } + for xname := range nextNodes { + if _, found := nextComponents[xname]; !found { + log.Error().Str("xname", xname).Msg("SMD component data missing node from Ethernet interface inventory") + return + } + } + 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 { @@ -344,6 +378,7 @@ func (s *SMDClient) PopulateNodes() { } s.nodes = nextNodes + s.components = nextComponents s.ipToXname = nextIPToXname s.macToXname = nextMACToXname s.wgipToXname = nextWGIPToXname @@ -431,6 +466,14 @@ func (s *SMDClient) ComponentInformation(id string) (base.Component, error) { if strings.Trim(id, " \t") == "" { return node, ErrEmptyID } + + s.nodesMutex.RLock() + if component, found := s.components[id]; found { + s.nodesMutex.RUnlock() + return cloneComponent(component), nil + } + s.nodesMutex.RUnlock() + ep := "/hsm/v2/State/Components/" + id err := s.getSMD(ep, &node) if err != nil { @@ -439,6 +482,14 @@ func (s *SMDClient) ComponentInformation(id string) (base.Component, error) { return node, nil } +func cloneComponent(component base.Component) base.Component { + if component.Enabled != nil { + enabled := *component.Enabled + component.Enabled = &enabled + } + return component +} + // ComponentInformationWithRetry wraps ComponentInformation with exponential backoff retry logic // for transient network errors and timeouts. This is critical during boot storms when SMD // may be temporarily overloaded. @@ -446,7 +497,7 @@ func (s *SMDClient) ComponentInformationWithRetry(id string, maxRetries int) (ba var lastErr error var node base.Component - for attempt := 0; attempt < maxRetries; attempt++ { + for attempt := range maxRetries { node, err := s.ComponentInformation(id) if err == nil { // Success - return immediately diff --git a/internal/smdclient/SMDclient_performance_test.go b/internal/smdclient/SMDclient_performance_test.go index 108d9081b..2a99c2751 100644 --- a/internal/smdclient/SMDclient_performance_test.go +++ b/internal/smdclient/SMDclient_performance_test.go @@ -19,6 +19,7 @@ func TestPopulateNodesBlockedRefreshDoesNotBlockCachedOperations(t *testing.T) { blockedPath string }{ {name: "blocked inventory request", blockedPath: "/hsm/v2/Inventory/EthernetInterfaces/"}, + {name: "blocked bulk component request", blockedPath: "/hsm/v2/State/Components"}, {name: "blocked bulk membership request", blockedPath: "/hsm/v2/memberships"}, } @@ -46,6 +47,8 @@ func TestPopulateNodesBlockedRefreshDoesNotBlockCachedOperations(t *testing.T) { t.Errorf("membership type query = %q, want node", got) } _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) } }) server := httptest.NewServer(handler) @@ -152,6 +155,8 @@ func TestGroupMembershipCached(t *testing.T) { ]`)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute","cabinet1"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) } }) server := httptest.NewServer(handler) @@ -167,11 +172,10 @@ func TestGroupMembershipCached(t *testing.T) { wgipToXname: make(map[string]string), } - // Populate cache - should make 2 requests (interfaces + membership) initialRequests := requestCount client.PopulateNodes() populateRequests := requestCount - initialRequests - assert.Equal(t, 2, populateRequests) + assert.Equal(t, 3, populateRequests) // Verify group membership was cached groups, err := client.GroupMembership("x1000") @@ -223,6 +227,8 @@ func TestConcurrentReads(t *testing.T) { {"id":"x1000","groupLabels":["compute"],"partitionName":""}, {"id":"x1001","groupLabels":["io"],"partitionName":""} ]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000", "x1001"}) } }) server := httptest.NewServer(handler) @@ -304,6 +310,8 @@ func TestReverseIndexPerformance(t *testing.T) { _, _ = w.Write([]byte(ethInterfaces)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(bulkMembershipsJSON(nodeCount))) + case "/hsm/v2/State/Components": + _, _ = w.Write([]byte(bulkComponentsJSON(nodeCount))) default: http.NotFound(w, r) } @@ -386,6 +394,8 @@ func TestCaseInsensitiveLookup(t *testing.T) { ]`)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) } }) server := httptest.NewServer(handler) @@ -447,6 +457,8 @@ func TestAddWGIPUpdatesReverseIndex(t *testing.T) { ]`)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) } }) server := httptest.NewServer(handler) @@ -508,6 +520,8 @@ func BenchmarkIDfromIP(b *testing.B) { _, _ = w.Write([]byte(ethInterfaces)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(bulkMembershipsJSON(1000))) + case "/hsm/v2/State/Components": + _, _ = w.Write([]byte(bulkComponentsJSON(1000))) default: http.NotFound(w, r) } @@ -552,6 +566,8 @@ func BenchmarkGroupMembership(b *testing.B) { ]`)) case "/hsm/v2/memberships": _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute","cabinet1","rack1"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) } }) server := httptest.NewServer(handler) @@ -585,3 +601,14 @@ func bulkMembershipsJSON(nodeCount int) string { } return memberships + "]" } + +func bulkComponentsJSON(nodeCount int) string { + components := `{"Components":[` + for i := range nodeCount { + if i > 0 { + components += "," + } + components += fmt.Sprintf(`{"ID":"x%d","Type":"Node","NID":"%d","Role":"compute"}`, i, i) + } + return components + "]}" +} diff --git a/internal/smdclient/SMDclient_stress_test.go b/internal/smdclient/SMDclient_stress_test.go new file mode 100644 index 000000000..b967f1ed2 --- /dev/null +++ b/internal/smdclient/SMDclient_stress_test.go @@ -0,0 +1,151 @@ +//go:build stress + +package smdclient + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + + base "github.com/Cray-HPE/hms-base" + "github.com/stretchr/testify/require" +) + +func TestStressComponentInformationCacheHits10K(t *testing.T) { + const componentCount = 10_000 + var liveLookups atomic.Int64 + client := &SMDClient{ + smdClient: &http.Client{Transport: failLiveComponentLookupRoundTripper{liveLookups: &liveLookups}}, + smdBaseURL: "http://smd.example", + nodesMutex: &sync.RWMutex{}, + components: make(map[string]base.Component, componentCount), + nodes: make(map[string]NodeMapping), + ipToXname: make(map[string]string), + macToXname: make(map[string]string), + wgipToXname: make(map[string]string), + accessToken: "token", + tokenEndpoint: "http://tokens.example", + } + for i := range componentCount { + id := fmt.Sprintf("x%d", i) + client.components[id] = base.Component{ID: id, Type: "Node", NID: jsonNumber(i), Role: "compute"} + } + + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(componentCount) + start.Add(1) + done.Add(componentCount) + + for i := range componentCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + + id := fmt.Sprintf("x%d", i) + component, err := client.ComponentInformationWithRetry(id, 3) + if err != nil { + t.Errorf("ComponentInformationWithRetry(%q) error = %v", id, err) + return + } + if component.ID != id || component.Role != "compute" { + t.Errorf("component = %+v, want ID %q role compute", component, id) + } + }() + } + + ready.Wait() + start.Done() + done.Wait() + require.Zero(t, liveLookups.Load()) +} + +func TestStressConcurrentGetSMDCoalescesTokenRefresh10K(t *testing.T) { + var tokenRequests atomic.Int64 + tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"fresh-token"}`)) + })) + defer tokenServer.Close() + + client := &SMDClient{ + smdClient: &http.Client{Transport: tokenAwareRoundTripper{}}, + smdBaseURL: "http://smd.example", + tokenEndpoint: tokenServer.URL, + accessToken: "stale-token", + nodesMutex: &sync.RWMutex{}, + } + + const requestCount = 10_000 + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(requestCount) + start.Add(1) + done.Add(requestCount) + + for range requestCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + + var response map[string]string + if err := client.getSMD("/component", &response); err != nil { + t.Errorf("getSMD() error = %v", err) + return + } + if response["ok"] != "true" { + t.Errorf("response = %v, want ok=true", response) + } + }() + } + + ready.Wait() + start.Done() + done.Wait() + + require.Equal(t, int64(1), tokenRequests.Load()) + require.Equal(t, "fresh-token", client.currentAccessToken()) +} + +type tokenAwareRoundTripper struct{} + +func (tokenAwareRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + statusCode := http.StatusOK + body := `{"ok":"true"}` + if req.Header.Get("Authorization") != "Bearer fresh-token" { + statusCode = http.StatusUnauthorized + body = `{"error":"unauthorized"}` + } + return &http.Response{ + StatusCode: statusCode, + Status: http.StatusText(statusCode), + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + Request: req, + }, nil +} + +type failLiveComponentLookupRoundTripper struct { + liveLookups *atomic.Int64 +} + +func (r failLiveComponentLookupRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + r.liveLookups.Add(1) + return nil, errors.New("unexpected live SMD request during cached component lookup") +} + +func jsonNumber(value int) json.Number { + return json.Number(fmt.Sprintf("%d", value)) +} diff --git a/internal/smdclient/SMDclient_test.go b/internal/smdclient/SMDclient_test.go index fa48ce9c8..4c6b93b73 100644 --- a/internal/smdclient/SMDclient_test.go +++ b/internal/smdclient/SMDclient_test.go @@ -2,6 +2,7 @@ package smdclient import ( "errors" + "fmt" "net/http" "net/http/httptest" "strings" @@ -10,6 +11,7 @@ import ( "testing" "time" + base "github.com/Cray-HPE/hms-base" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -69,6 +71,8 @@ func TestPopulateNodes(t *testing.T) { {"id":"x1003","groupLabels":["compute","cabinet1"],"partitionName":""}, {"id":"x9999","groupLabels":["unrelated"],"partitionName":""} ]`)) + case "/hsm/v2/State/Components": + writeTestComponents(w) } }) server := httptest.NewServer(handler) @@ -127,12 +131,136 @@ func TestPopulateNodes(t *testing.T) { defer requestsMutex.Unlock() require.Equal(t, []string{ "/hsm/v2/Inventory/EthernetInterfaces/", + "/hsm/v2/State/Components", "/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 TestConcurrentGetSMDCoalescesTokenRefresh(t *testing.T) { + var tokenRequests atomic.Int64 + tokenServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tokenRequests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"fresh-token"}`)) + })) + defer tokenServer.Close() + + smdServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bearer fresh-token" { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":"unauthorized"}`)) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":"true"}`)) + })) + defer smdServer.Close() + + client := &SMDClient{ + smdClient: smdServer.Client(), + smdBaseURL: smdServer.URL, + tokenEndpoint: tokenServer.URL, + accessToken: "stale-token", + nodesMutex: &sync.RWMutex{}, + } + + const requestCount = 64 + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(requestCount) + start.Add(1) + done.Add(requestCount) + for range requestCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + var response map[string]string + if err := client.getSMD("/component", &response); err != nil { + t.Errorf("getSMD() error = %v", err) + return + } + if response["ok"] != "true" { + t.Errorf("response = %v, want ok=true", response) + } + }() + } + ready.Wait() + start.Done() + done.Wait() + + if got := tokenRequests.Load(); got != 1 { + t.Fatalf("token endpoint requests = %d, want 1", got) + } + if got := client.currentAccessToken(); got != "fresh-token" { + t.Fatalf("current access token = %q, want fresh-token", got) + } +} + +func TestComponentInformationUsesCache(t *testing.T) { + var perNodeComponentRequests atomic.Int64 + server := httptest.NewServer(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/State/Components": + writeComponents(w, []string{"x1000"}) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + case "/hsm/v2/State/Components/x1000": + perNodeComponentRequests.Add(1) + w.WriteHeader(http.StatusInternalServerError) + default: + w.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestSMDClient(server) + client.PopulateNodes() + + component, err := client.ComponentInformation("x1000") + require.NoError(t, err) + require.Equal(t, "x1000", component.ID) + require.Equal(t, "compute", component.Role) + require.Zero(t, perNodeComponentRequests.Load()) + + component.Role = "mutated" + component.Enabled = boolPtr(false) + fresh, err := client.ComponentInformationWithRetry("x1000", 3) + require.NoError(t, err) + require.Equal(t, "compute", fresh.Role) + require.Nil(t, fresh.Enabled) + require.Zero(t, perNodeComponentRequests.Load()) +} + +func TestComponentInformationFallsBackOnCacheMiss(t *testing.T) { + var perNodeComponentRequests atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/hsm/v2/State/Components/x9999": + perNodeComponentRequests.Add(1) + _, _ = w.Write([]byte(`{"ID":"x9999","Type":"Node","NID":"9999","Role":"fallback"}`)) + default: + w.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestSMDClient(server) + component, err := client.ComponentInformation("x9999") + require.NoError(t, err) + require.Equal(t, "x9999", component.ID) + require.Equal(t, "fallback", component.Role) + require.Equal(t, int64(1), perNodeComponentRequests.Load()) +} + func TestIPfromID(t *testing.T) { // Mock SMD server handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -180,6 +308,8 @@ func TestIPfromID(t *testing.T) { {"id":"x1002","groupLabels":["compute"],"partitionName":""}, {"id":"x1003","groupLabels":["compute"],"partitionName":""} ]`)) + case "/hsm/v2/State/Components": + writeTestComponents(w) } }) server := httptest.NewServer(handler) @@ -272,6 +402,8 @@ func TestIDfromIP(t *testing.T) { {"id":"x1002","groupLabels":["compute"],"partitionName":""}, {"id":"x1003","groupLabels":["compute"],"partitionName":""} ]`)) + case "/hsm/v2/State/Components": + writeTestComponents(w) } }) server := httptest.NewServer(handler) @@ -364,6 +496,8 @@ func TestIDfromMAC(t *testing.T) { {"id":"x1002","groupLabels":["compute"],"partitionName":""}, {"id":"x1003","groupLabels":["compute"],"partitionName":""} ]`)) + case "/hsm/v2/State/Components": + writeTestComponents(w) } }) server := httptest.NewServer(handler) @@ -424,6 +558,8 @@ func TestPopulateNodesMissingMembershipUsesEmptyGroups(t *testing.T) { t.Errorf("membership type query = %q, want node", got) } _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":"ignored"}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000", "x1001"}) default: t.Errorf("unexpected SMD request: %s", r.URL.RequestURI()) w.WriteHeader(http.StatusNotFound) @@ -476,6 +612,8 @@ func TestPopulateNodesBulkMembershipFailurePreservesCache(t *testing.T) { return } _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + case "/hsm/v2/State/Components": + writeComponents(w, []string{"x1000"}) default: if strings.HasPrefix(r.URL.Path, "/hsm/v2/memberships/") { perNodeRequests.Add(1) @@ -513,6 +651,76 @@ func TestPopulateNodesBulkMembershipFailurePreservesCache(t *testing.T) { } } +func TestPopulateNodesBulkComponentFailurePreservesCache(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(`{"Components":[{"ID":"x1000"}`)) + }, + }, + { + name: "missing component", + writeFailure: func(w http.ResponseWriter) { + _, _ = w.Write([]byte(`{"Components":[]}`)) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var failComponents atomic.Bool + server := httptest.NewServer(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/State/Components": + if failComponents.Load() { + tt.writeFailure(w) + return + } + writeComponents(w, []string{"x1000"}) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + default: + w.WriteHeader(http.StatusNotFound) + } + })) + defer server.Close() + + client := newTestSMDClient(server) + client.PopulateNodes() + component, err := client.ComponentInformation("x1000") + require.NoError(t, err) + require.Equal(t, "compute", component.Role) + + client.nodesMutex.RLock() + oldTimestamp := client.nodes_last_update + client.nodesMutex.RUnlock() + failComponents.Store(true) + client.PopulateNodes() + + fresh, err := client.ComponentInformation("x1000") + require.NoError(t, err) + require.Equal(t, "compute", fresh.Role) + client.nodesMutex.RLock() + require.Equal(t, oldTimestamp, client.nodes_last_update) + client.nodesMutex.RUnlock() + }) + } +} + func newTestSMDClient(server *httptest.Server) *SMDClient { return &SMDClient{ smdClient: server.Client(), @@ -522,9 +730,29 @@ func newTestSMDClient(server *httptest.Server) *SMDClient { ipToXname: make(map[string]string), macToXname: make(map[string]string), wgipToXname: make(map[string]string), + components: make(map[string]base.Component), } } +func writeTestComponents(w http.ResponseWriter) { + writeComponents(w, []string{"x1000", "x1001", "x1002", "x1003"}) +} + +func writeComponents(w http.ResponseWriter, ids []string) { + _, _ = w.Write([]byte(`{"Components":[`)) + for i, id := range ids { + if i > 0 { + _, _ = w.Write([]byte(`,`)) + } + _, _ = fmt.Fprintf(w, `{"ID":%q,"Type":"Node","NID":%q,"Role":"compute"}`, id, strings.TrimPrefix(id, "x")) + } + _, _ = w.Write([]byte(`]}`)) +} + +func boolPtr(value bool) *bool { + return &value +} + func mustIDfromIP(t *testing.T, client *SMDClient, ip string) string { t.Helper() id, err := client.IDfromIP(ip) diff --git a/internal/smdclient/oidc.go b/internal/smdclient/oidc.go index b044b53b4..dfc77105b 100644 --- a/internal/smdclient/oidc.go +++ b/internal/smdclient/oidc.go @@ -19,11 +19,27 @@ type oidcTokenData struct { // authorization grant. Support for said grant should probably be implemented // at some point. func (s *SMDClient) RefreshToken() error { + s.accessTokenMutex.Lock() + defer s.accessTokenMutex.Unlock() + return s.refreshTokenLocked() +} + +func (s *SMDClient) refreshTokenIfCurrent(rejectedToken string) error { + s.accessTokenMutex.Lock() + defer s.accessTokenMutex.Unlock() + if s.accessToken != rejectedToken { + return nil + } + return s.refreshTokenLocked() +} + +func (s *SMDClient) refreshTokenLocked() error { // Request new token from OIDC server r, err := http.Get(s.tokenEndpoint) if err != nil { return err } + defer r.Body.Close() body, err := io.ReadAll(r.Body) if err != nil { return err diff --git a/pkg/wgtunnel/allocator.go b/pkg/wgtunnel/allocator.go index 4d22681f9..4020be505 100644 --- a/pkg/wgtunnel/allocator.go +++ b/pkg/wgtunnel/allocator.go @@ -13,6 +13,7 @@ type IPAllocator struct { mu sync.Mutex networkAddr net.IP broadcastAddr net.IP + nextIP net.IP } // NewIPAllocator initializes a new IPAllocator for a given network. @@ -39,6 +40,7 @@ func NewIPAllocator(cidr string) (*IPAllocator, error) { networkAddr: networkAddr, broadcastAddr: broadcastAddr, usedIPs: make(map[string]bool), + nextIP: nextIP(networkAddr), }, nil } @@ -56,6 +58,9 @@ func (a *IPAllocator) Reserve(ipAddr net.IPAddr) error { return errors.New("IP address already allocated") } a.usedIPs[ipStr] = true + if ip.Equal(a.nextIP) { + a.nextIP = nextIP(ip) + } return nil } @@ -64,26 +69,27 @@ func (a *IPAllocator) NextAvailable() (net.IPAddr, error) { a.mu.Lock() defer a.mu.Unlock() - ip := make(net.IP, len(a.networkAddr)) - copy(ip, a.networkAddr) + ip := cloneIP(a.nextIP) + start := cloneIP(a.nextIP) + wrapped := false for { - for i := len(ip) - 1; i >= 0; i-- { - ip[i]++ - if ip[i] != 0 { - break - } + if wrapped && ip.Equal(start) { + return net.IPAddr{}, errors.New("IP range exhausted: no available IP addresses in range " + a.network.String()) } - // Check if the incremented IP is still within the subnet range - if !a.network.Contains(ip) { - return net.IPAddr{}, errors.New("IP range exhausted: no available IP addresses in range " + a.network.String()) + if !a.network.Contains(ip) || ip.Equal(a.networkAddr) || ip.Equal(a.broadcastAddr) { + ip = a.firstUsableIP() + wrapped = true + continue } ipStr := ip.String() if !a.usedIPs[ipStr] { a.usedIPs[ipStr] = true + a.nextIP = nextIP(ip) return net.IPAddr{IP: ip}, nil } + ip = nextIP(ip) } } @@ -105,5 +111,46 @@ func (a *IPAllocator) Release(ipAddr net.IPAddr) error { return errors.New("IP address not allocated") } delete(a.usedIPs, ipStr) + if a.network.Contains(ipAddr.IP) && !ipAddr.IP.Equal(a.networkAddr) && !ipAddr.IP.Equal(a.broadcastAddr) && ipLess(ipAddr.IP, a.nextIP) { + a.nextIP = cloneIP(ipAddr.IP) + } return nil } + +func (a *IPAllocator) firstUsableIP() net.IP { + return nextIP(a.networkAddr) +} + +func cloneIP(ip net.IP) net.IP { + clone := make(net.IP, len(ip)) + copy(clone, ip) + return clone +} + +func nextIP(ip net.IP) net.IP { + next := cloneIP(ip) + for i := len(next) - 1; i >= 0; i-- { + next[i]++ + if next[i] != 0 { + break + } + } + return next +} + +func ipLess(left, right net.IP) bool { + left = left.To4() + right = right.To4() + if left == nil || right == nil { + return false + } + for i := range left { + if left[i] < right[i] { + return true + } + if left[i] > right[i] { + return false + } + } + return false +} diff --git a/pkg/wgtunnel/allocator_test.go b/pkg/wgtunnel/allocator_test.go index 4f33353bd..c75cc4beb 100644 --- a/pkg/wgtunnel/allocator_test.go +++ b/pkg/wgtunnel/allocator_test.go @@ -28,6 +28,22 @@ func TestReserve(t *testing.T) { } } +func TestReserveAdvancesCursorWhenReservingNextIP(t *testing.T) { + allocator, _ := NewIPAllocator("192.168.1.0/24") // ignoring error on NewIPAllocator. We'll catch it on use anyway. + if err := allocator.Reserve(net.IPAddr{IP: net.ParseIP("192.168.1.1")}); err != nil { + t.Fatalf("Failed to reserve IP: %v", err) + } + + ip, err := allocator.NextAvailable() + if err != nil { + t.Fatalf("Failed to get next available IP: %v", err) + } + expectedIP := net.IPAddr{IP: net.ParseIP("192.168.1.2")} + if !ip.IP.Equal(expectedIP.IP) { + t.Fatalf("Expected IP %v, got %v", expectedIP, ip) + } +} + func TestNextAvailable(t *testing.T) { allocator, _ := NewIPAllocator("192.168.1.0/24") // ignoring error on NewIPAllocator. We'll catch it on use anyway. @@ -52,6 +68,10 @@ func TestNextAvailable(t *testing.T) { if allocator.IsAllocated(broadcastIP) { t.Fatalf("Broadcast address should not be allocated") } + _, err := allocator.NextAvailable() + if err == nil { + t.Fatalf("Expected error when IP range is exhausted") + } } func TestIsAllocated(t *testing.T) { @@ -89,3 +109,25 @@ func TestRelease(t *testing.T) { t.Fatalf("Expected error when releasing a non-allocated IP") } } + +func TestReleaseMakesLowerIPAvailableAgain(t *testing.T) { + allocator, _ := NewIPAllocator("192.168.1.0/24") + for i := 1; i <= 10; i++ { + _, err := allocator.NextAvailable() + if err != nil { + t.Fatalf("Failed to get next available IP: %v", err) + } + } + releasedIP := net.IPAddr{IP: net.ParseIP("192.168.1.5")} + if err := allocator.Release(releasedIP); err != nil { + t.Fatalf("Failed to release IP: %v", err) + } + + ip, err := allocator.NextAvailable() + if err != nil { + t.Fatalf("Failed to get next available IP: %v", err) + } + if !ip.IP.Equal(releasedIP.IP) { + t.Fatalf("Expected released IP %v, got %v", releasedIP, ip) + } +} diff --git a/pkg/wgtunnel/handlers.go b/pkg/wgtunnel/handlers.go index 0d32000c8..e308216b7 100644 --- a/pkg/wgtunnel/handlers.go +++ b/pkg/wgtunnel/handlers.go @@ -107,7 +107,7 @@ func AddClientHandler(im *InterfaceManager, smdClient smdclient.SMDClientInterfa // Add the client to the WireGuard configuration. log.Info().Msgf("Adding WireGuard peer: PublicKey=%s, ClientVPNIP=%s, ClientIP=%s\n", publicKey, clientVPNIP, clientIP) - if err := im.AddPeer(im.GetInterfaceName(), publicKey, clientVPNIP, clientIP); err != nil { + if err := im.AddPeer(publicKey, clientVPNIP, clientIP); err != nil { http.Error(w, "Failed to configure WireGuard tunnel: "+err.Error(), http.StatusInternalServerError) return } diff --git a/pkg/wgtunnel/tunnels.go b/pkg/wgtunnel/tunnels.go index 570bb016f..f2197fe5c 100644 --- a/pkg/wgtunnel/tunnels.go +++ b/pkg/wgtunnel/tunnels.go @@ -148,13 +148,29 @@ func (m *InterfaceManager) IpForPeer(peerName string, publicKey string) string { } func (m *InterfaceManager) RemovePeer(peerName string) error { - m.peersMutex.Lock() - defer m.peersMutex.Unlock() - if err := exec.Command("wg", "set", m.interfaceName, "peer", m.peers[peerName].PublicKey, "remove").Run(); err != nil { + m.peersMutex.RLock() + peer, found := m.peers[peerName] + interfaceName := m.interfaceName + m.peersMutex.RUnlock() + if !found { + return nil + } + + if err := exec.Command("wg", "set", interfaceName, "peer", peer.PublicKey, "remove").Run(); err != nil { log.Error().Err(err).Msgf("Failed to remove peer (%s)", peerName) return err } - delete(m.peers, peerName) + + m.peersMutex.Lock() + defer m.peersMutex.Unlock() + currentPeer, found := m.peers[peerName] + if found && currentPeer.PublicKey == peer.PublicKey { + delete(m.peers, peerName) + if err := m.ipManager.Release(peer.IP); err != nil { + log.Error().Err(err).Msgf("Failed to release peer IP (%s)", peerName) + return err + } + } return nil } @@ -251,19 +267,8 @@ func (m *InterfaceManager) StopServer() error { return nil } -func (m *InterfaceManager) AddPeer(peerName, publicKey, vpnIP, clientIP string) error { - m.peersMutex.RLock() - defer m.peersMutex.RUnlock() - - // Add the peer to the WireGuard configuration - if err := AddWireGuardPeer(m.interfaceName, publicKey, vpnIP, clientIP); err != nil { - return err - } - m.peers[peerName] = PeerConfig{ - PublicKey: publicKey, - IP: net.IPAddr{IP: net.ParseIP(vpnIP), Zone: ""}, - } - return nil +func (m *InterfaceManager) AddPeer(publicKey, vpnIP, clientIP string) error { + return AddWireGuardPeer(m.interfaceName, publicKey, vpnIP, clientIP) } // AddWireGuardPeer adds a peer to the WireGuard configuration. diff --git a/pkg/wgtunnel/tunnels_stress_test.go b/pkg/wgtunnel/tunnels_stress_test.go new file mode 100644 index 000000000..bc5bce939 --- /dev/null +++ b/pkg/wgtunnel/tunnels_stress_test.go @@ -0,0 +1,52 @@ +//go:build stress + +package wgtunnel + +import ( + "fmt" + "sync" + "testing" +) + +func TestStressIpForPeerConcurrent10K(t *testing.T) { + manager := newTestInterfaceManager(t) + + const peerCount = 10_000 + var ready sync.WaitGroup + var start sync.WaitGroup + var done sync.WaitGroup + ready.Add(peerCount) + start.Add(1) + done.Add(peerCount) + + allocated := make(chan string, peerCount) + for i := range peerCount { + go func() { + defer done.Done() + ready.Done() + start.Wait() + + peerName := fmt.Sprintf("10.%d.%d.%d", i/65536, (i/256)%256, i%256) + allocated <- manager.IpForPeer(peerName, fmt.Sprintf("key-%d", i)) + }() + } + + ready.Wait() + start.Done() + done.Wait() + close(allocated) + + seen := make(map[string]struct{}, peerCount) + for ip := range allocated { + if ip == "" { + t.Fatal("expected allocated IP, got empty string") + } + if _, found := seen[ip]; found { + t.Fatalf("IP %s allocated more than once", ip) + } + seen[ip] = struct{}{} + } + if len(seen) != peerCount { + t.Fatalf("expected %d allocations, got %d", peerCount, len(seen)) + } +} diff --git a/pkg/wgtunnel/tunnels_test.go b/pkg/wgtunnel/tunnels_test.go index d1a203eac..97fc46f81 100644 --- a/pkg/wgtunnel/tunnels_test.go +++ b/pkg/wgtunnel/tunnels_test.go @@ -3,8 +3,12 @@ package wgtunnel import ( "fmt" "net" + "os" + "path/filepath" + "strings" "sync" "testing" + "time" ) func newTestInterfaceManager(t *testing.T) *InterfaceManager { @@ -94,3 +98,154 @@ func TestGetPeersReturnsCopy(t *testing.T) { t.Fatal("GetPeers returned mutable internal peers map") } } +func TestAddPeerConfiguresWireGuardWithoutMutatingPeers(t *testing.T) { + manager := newTestInterfaceManager(t) + clientIP := "10.1.0.1" + publicKey := "key-1" + vpnIP := manager.IpForPeer(clientIP, publicKey) + if vpnIP == "" { + t.Fatal("expected allocated peer IP") + } + + argsFile := installFakeWG(t) + if err := manager.AddPeer(publicKey, vpnIP, clientIP); err != nil { + t.Fatalf("AddPeer() error = %v, want nil", err) + } + + manager.peersMutex.RLock() + defer manager.peersMutex.RUnlock() + if _, found := manager.peers[manager.GetInterfaceName()]; found { + t.Fatalf("AddPeer wrote peer under interface name %q", manager.GetInterfaceName()) + } + peer, found := manager.peers[clientIP] + if !found { + t.Fatalf("peer %q missing after IpForPeer", clientIP) + } + if peer.PublicKey != publicKey || peer.IP.IP.String() != vpnIP { + t.Fatalf("peer = %+v, want public key %q and IP %q", peer, publicKey, vpnIP) + } + + args, err := os.ReadFile(argsFile) + if err != nil { + t.Fatalf("failed to read fake wg args: %v", err) + } + got := strings.TrimSpace(string(args)) + want := fmt.Sprintf("set wg0 peer %s allowed-ips %s/32", publicKey, vpnIP) + if got != want { + t.Fatalf("wg args = %q, want %q", got, want) + } +} + +func TestRemovePeerRunsWireGuardOutsidePeerLock(t *testing.T) { + manager := newTestInterfaceManager(t) + peerName := "10.1.0.1" + originalKey := "key-1" + vpnIP := manager.IpForPeer(peerName, originalKey) + if vpnIP == "" { + t.Fatal("expected allocated peer IP") + } + + release := make(chan struct{}) + argsFile := installBlockingFakeWG(t, release) + removeDone := make(chan error, 1) + go func() { + removeDone <- manager.RemovePeer(peerName) + }() + + waitForWGArgs(t, argsFile) + if got := manager.IpForPeer(peerName, "key-2"); got != vpnIP { + t.Fatalf("IpForPeer while remove was blocked = %q, want %q", got, vpnIP) + } + close(release) + if err := <-removeDone; err != nil { + t.Fatalf("RemovePeer() error = %v", err) + } + + manager.peersMutex.RLock() + defer manager.peersMutex.RUnlock() + peer, found := manager.peers[peerName] + if !found { + t.Fatal("RemovePeer deleted peer that was replaced while wg command ran") + } + if peer.PublicKey != "key-2" { + t.Fatalf("peer public key = %q, want replacement key", peer.PublicKey) + } + + args, err := os.ReadFile(argsFile) + if err != nil { + t.Fatalf("failed to read fake wg args: %v", err) + } + want := fmt.Sprintf("set wg0 peer %s remove", originalKey) + if got := strings.TrimSpace(string(args)); got != want { + t.Fatalf("wg args = %q, want %q", got, want) + } +} + +func TestRemovePeerReleasesAllocatedIP(t *testing.T) { + manager := newTestInterfaceManager(t) + peerName := "10.1.0.1" + publicKey := "key-1" + vpnIP := manager.IpForPeer(peerName, publicKey) + if vpnIP == "" { + t.Fatal("expected allocated peer IP") + } + + installFakeWG(t) + if err := manager.RemovePeer(peerName); err != nil { + t.Fatalf("RemovePeer() error = %v, want nil", err) + } + + reusedIP := manager.IpForPeer("10.1.0.2", "key-2") + if reusedIP != vpnIP { + t.Fatalf("reused IP = %q, want released IP %q", reusedIP, vpnIP) + } +} + +func installFakeWG(t *testing.T) string { + t.Helper() + + dir := t.TempDir() + argsFile := filepath.Join(dir, "wg-args") + wgPath := filepath.Join(dir, "wg") + script := fmt.Sprintf("#!/bin/sh\nprintf '%%s\\n' \"$*\" >> %q\n", argsFile) + if err := os.WriteFile(wgPath, []byte(script), 0o755); err != nil { + t.Fatalf("failed to write fake wg: %v", err) + } + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + return argsFile +} + +func installBlockingFakeWG(t *testing.T, release <-chan struct{}) string { + t.Helper() + + dir := t.TempDir() + argsFile := filepath.Join(dir, "wg-args") + releaseFile := filepath.Join(dir, "release") + wgPath := filepath.Join(dir, "wg") + script := fmt.Sprintf("#!/bin/sh\nprintf '%%s\\n' \"$*\" >> %q\nwhile [ ! -f %q ]; do sleep 0.01; done\n", argsFile, releaseFile) + if err := os.WriteFile(wgPath, []byte(script), 0o755); err != nil { + t.Fatalf("failed to write fake wg: %v", err) + } + t.Setenv("PATH", dir+string(os.PathListSeparator)+os.Getenv("PATH")) + go func() { + <-release + _ = os.WriteFile(releaseFile, []byte("release"), 0o644) + }() + return argsFile +} + +func waitForWGArgs(t *testing.T, argsFile string) { + t.Helper() + deadline := time.After(time.Second) + for { + select { + case <-deadline: + t.Fatal("fake wg command did not start") + default: + if _, err := os.Stat(argsFile); err == nil { + return + } + time.Sleep(10 * time.Millisecond) + } + } +}