diff --git a/pkg/device/cambricon/device.go b/pkg/device/cambricon/device.go index 9e3defd85b..e30e38e749 100644 --- a/pkg/device/cambricon/device.go +++ b/pkg/device/cambricon/device.go @@ -18,18 +18,16 @@ package cambricon import ( "context" - "encoding/json" "flag" "fmt" - "math/rand" "slices" "strings" - "time" "github.com/Project-HAMi/HAMi/pkg/device" "github.com/Project-HAMi/HAMi/pkg/device/common" "github.com/Project-HAMi/HAMi/pkg/util" "github.com/Project-HAMi/HAMi/pkg/util/client" + "github.com/Project-HAMi/HAMi/pkg/util/nodelock" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" @@ -46,14 +44,11 @@ const ( MluMemSplitEnable = "CAMBRICON_SPLIT_ENABLE" MLUInUse = "cambricon.com/use-mlutype" MLUNoUse = "cambricon.com/nouse-mlutype" - // MLUUseUUID annotation specifies a comma-separated list of MLU UUIDs to use. - MLUUseUUID = "cambricon.com/use-gpuuuid" - // MLUNoUseUUID annotation specifies a comma-separated list of MLU UUIDs to exclude. - MLUNoUseUUID = "cambricon.com/nouse-gpuuuid" - DsmluLockTime = "cambricon.com/dsmlu.lock" + MLUUseUUID = "cambricon.com/use-gpuuuid" + MLUNoUseUUID = "cambricon.com/nouse-gpuuuid" DsmluProfile = "CAMBRICON_DSMLU_PROFILE" DsmluResourceAssigned = "CAMBRICON_DSMLU_ASSIGNED" - retry = 5 + dsmluLockTime = "cambricon.com/dsmlu.lock" ) var ( @@ -93,88 +88,46 @@ func (dev *CambriconDevices) CommonWord() string { return CambriconMLUCommonWord } -func (dev *CambriconDevices) setNodeLock(node *corev1.Node) error { - ctx := context.Background() - if _, ok := node.Annotations[DsmluLockTime]; ok { - return fmt.Errorf("node %s is locked", node.Name) - } - - patchedAnnotation, err := json.Marshal( - map[string]any{ - "metadata": map[string]map[string]string{"annotations": { - DsmluLockTime: time.Now().Format(time.RFC3339), - }}}) - if err != nil { - klog.ErrorS(err, "Failed to patch node annotation", "node", node.Name) - return fmt.Errorf("patch node annotation %v", err) - } - - _, err = client.GetClient().CoreV1().Nodes().Patch(ctx, node.Name, types.StrategicMergePatchType, patchedAnnotation, metav1.PatchOptions{}) - for i := 0; i < retry && err != nil; i++ { - klog.ErrorS(err, "Failed to patch node annotation", "node", node.Name, "retry", i) - time.Sleep(time.Duration(rand.Intn(i+1)) * 10 * time.Millisecond) - _, err = client.GetClient().CoreV1().Nodes().Patch(ctx, node.Name, types.StrategicMergePatchType, patchedAnnotation, metav1.PatchOptions{}) - } - if err != nil { - return fmt.Errorf("setNodeLock exceeds retry count %d", retry) - } - klog.InfoS("Node lock set", "node", node.Name) - return nil -} - -func (dev *CambriconDevices) LockNode(n *corev1.Node, p *corev1.Pod) error { - found := false +func (dev *CambriconDevices) hasMLURequest(p *corev1.Pod) bool { for _, val := range p.Spec.Containers { if (dev.GenerateResourceRequests(&val).Nums) > 0 { - found = true - break + return true } } - if !found { + return false +} + +func (dev *CambriconDevices) LockNode(n *corev1.Node, p *corev1.Pod) error { + if !dev.hasMLURequest(p) { return nil } - if _, ok := n.Annotations[DsmluLockTime]; !ok { - return dev.setNodeLock(n) - } - lockTime, err := time.Parse(time.RFC3339, n.Annotations[DsmluLockTime]) - if err != nil { - return err - } - if time.Since(lockTime) > time.Minute*2 { - klog.InfoS("Node lock expired", "node", n.Name, "lockTime", lockTime) - err = dev.ReleaseNodeLock(n, p) - if err != nil { - klog.ErrorS(err, "Failed to release node lock", "node", n.Name) - return err - } - return dev.setNodeLock(n) - } - return fmt.Errorf("node %s has been locked within 2 minutes", n.Name) + dev.cleanupLegacyLock(n.Name) + return nodelock.LockNode(n.Name, nodelock.NodeLockKey, p) } func (dev *CambriconDevices) ReleaseNodeLock(n *corev1.Node, p *corev1.Pod) error { - if n.Annotations == nil { - return nil - } - if _, ok := n.Annotations[DsmluLockTime]; !ok { - klog.InfoS("Node lock not set", "node", n.Name) + if !dev.hasMLURequest(p) { return nil } + return nodelock.ReleaseNodeLock(n.Name, nodelock.NodeLockKey, p, false) +} - newNode := n.DeepCopy() - delete(newNode.Annotations, DsmluLockTime) - _, err := client.GetClient().CoreV1().Nodes().Update(context.Background(), newNode, metav1.UpdateOptions{}) - for i := 0; i < retry && err != nil; i++ { - klog.ErrorS(err, "Failed to patch node annotation", "node", n.Name, "retry", i) - time.Sleep(time.Duration(rand.Intn(i+1)) * 10 * time.Millisecond) - _, err = client.GetClient().CoreV1().Nodes().Update(context.Background(), newNode, metav1.UpdateOptions{}) +func (dev *CambriconDevices) cleanupLegacyLock(nodeName string) { + node, err := client.GetClient().CoreV1().Nodes().Get(context.Background(), nodeName, metav1.GetOptions{}) + if err != nil { + klog.V(4).InfoS("cleanupLegacyLock: failed to get node", "node", nodeName, "err", err) + return + } + if _, ok := node.Annotations[dsmluLockTime]; !ok { + return } + patch := []byte(`[{"op":"remove","path":"/metadata/annotations/cambricon.com~1dsmlu.lock"}]`) + _, err = client.GetClient().CoreV1().Nodes().Patch(context.Background(), nodeName, types.JSONPatchType, patch, metav1.PatchOptions{}) if err != nil { - return fmt.Errorf("releaseNodeLock exceeds retry count %d", retry) + klog.V(4).InfoS("cleanupLegacyLock: failed to remove legacy annotation", "node", nodeName, "err", err) + return } - delete(n.Annotations, DsmluLockTime) - klog.InfoS("Node lock released", "node", n.Name) - return nil + klog.InfoS("cleanupLegacyLock: removed legacy lock annotation", "node", nodeName) } func (dev *CambriconDevices) NodeCleanUp(nn string) error { diff --git a/pkg/device/cambricon/device_test.go b/pkg/device/cambricon/device_test.go index 3c38406049..f77c6e5bff 100644 --- a/pkg/device/cambricon/device_test.go +++ b/pkg/device/cambricon/device_test.go @@ -19,10 +19,9 @@ package cambricon import ( "context" "flag" - "strings" "testing" - "time" + "github.com/Project-HAMi/HAMi/pkg/util/nodelock" "github.com/stretchr/testify/assert" corev1 "k8s.io/api/core/v1" "k8s.io/apimachinery/pkg/api/resource" @@ -394,110 +393,98 @@ func Test_PatchAnnotations(t *testing.T) { } } -func Test_setNodeLock(t *testing.T) { +func TestLockNode(t *testing.T) { + config := CambriconConfig{ + ResourceCountName: MLUResourceCount, + ResourceMemoryName: MLUResourceMemory, + ResourceCoreName: MLUResourceCores, + } + tests := []struct { - name string - node corev1.Node - expectErr bool - expectMsg string + name string + pod *corev1.Pod + hasLock bool }{ { - name: "node is locked", - node: corev1.Node{ - ObjectMeta: metav1.ObjectMeta{ - Name: "node-01", - Annotations: map[string]string{ - "cambricon.com/dsmlu.lock": "test123", - }, + name: "no MLU containers skip lock", + pod: &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{Name: "pod-no-mlu", Namespace: "default"}, + Spec: corev1.PodSpec{Containers: []corev1.Container{{Name: "app"}}}, + }, + hasLock: false, + }, + { + name: "has MLU container acquires lock", + pod: &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{Name: "pod-mlu", Namespace: "default"}, + Spec: corev1.PodSpec{ + Containers: []corev1.Container{{ + Name: "mlu-app", + Resources: corev1.ResourceRequirements{ + Limits: corev1.ResourceList{ + corev1.ResourceName(MLUResourceCount): *resource.NewQuantity(1, resource.BinarySI), + corev1.ResourceName(MLUResourceMemory): resource.MustParse("2048"), + corev1.ResourceName(MLUResourceCores): resource.MustParse("100"), + }, + }, + }}, }, }, - expectErr: true, - expectMsg: "node node-01 is locked", + hasLock: true, }, - { - name: "set node lock", - node: corev1.Node{ + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + client.KubeClient = fake.NewClientset() + node := &corev1.Node{ ObjectMeta: metav1.ObjectMeta{ - Name: "node-02", + Name: "test-node", Annotations: map[string]string{}, }, - }, - expectErr: false, - }, - } - - client.KubeClient = fake.NewClientset() - k8sClient := client.GetClient() - if k8sClient != nil { - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ctx := context.Background() - - defer func() { - if tt.node.Name != "" { - // Delete the node to clean up - err := k8sClient.CoreV1().Nodes().Delete(ctx, tt.node.Name, metav1.DeleteOptions{}) - if err != nil { - t.Errorf("failed to delete node %s: %v", tt.node.Name, err) - } - } - }() - - _, err := k8sClient.CoreV1().Nodes().Create(ctx, &tt.node, metav1.CreateOptions{}) - if err != nil { - t.Fatalf("failed to create node %s: %v", tt.node.Name, err) - } + } + _, err := client.KubeClient.CoreV1().Nodes().Create(context.Background(), node, metav1.CreateOptions{}) + assert.NoError(t, err) - dev := CambriconDevices{} - err = dev.setNodeLock(&tt.node) + dev := InitMLUDevice(config) + err = dev.LockNode(node, tt.pod) + assert.NoError(t, err) - if tt.expectErr { - if err == nil { - t.Errorf("expected error but got none") - } else if !strings.Contains(err.Error(), tt.expectMsg) { - t.Errorf("expected error to contain '%s' but got '%s'", tt.expectMsg, err.Error()) - } - } else { - if err != nil { - t.Errorf("did not expect error but got %v", err) - } - } - }) - } + updated, err := client.KubeClient.CoreV1().Nodes().Get(context.Background(), "test-node", metav1.GetOptions{}) + assert.NoError(t, err) + _, ok := updated.Annotations[nodelock.NodeLockKey] + assert.Equal(t, tt.hasLock, ok) + }) } } -// Setup function to initialize resources for each test case. -func setupTest(t *testing.T) (*corev1.Node, *corev1.Pod, func(), *fake.Clientset) { - ctx := context.Background() - - clientset := fake.NewClientset() +func TestReleaseNodeLock(t *testing.T) { + config := CambriconConfig{ + ResourceCountName: MLUResourceCount, + ResourceMemoryName: MLUResourceMemory, + ResourceCoreName: MLUResourceCores, + } + client.KubeClient = fake.NewClientset() node := &corev1.Node{ ObjectMeta: metav1.ObjectMeta{ Name: "test-node", - }, - Status: corev1.NodeStatus{ - Capacity: corev1.ResourceList{ - corev1.ResourceName(MLUResourceCount): resource.MustParse("2"), - corev1.ResourceName(MLUResourceMemory): resource.MustParse("4096"), - corev1.ResourceName(MLUResourceCores): resource.MustParse("200"), + Annotations: map[string]string{ + nodelock.NodeLockKey: "test-lock", }, }, } + _, err := client.KubeClient.CoreV1().Nodes().Create(context.Background(), node, metav1.CreateOptions{}) + assert.NoError(t, err) + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{Name: "pod-mlu", Namespace: "default"}, Spec: corev1.PodSpec{ Containers: []corev1.Container{{ - Name: "test-container", + Name: "mlu-app", Resources: corev1.ResourceRequirements{ Limits: corev1.ResourceList{ - corev1.ResourceName(MLUResourceCount): resource.MustParse("1"), - corev1.ResourceName(MLUResourceMemory): resource.MustParse("2048"), - corev1.ResourceName(MLUResourceCores): resource.MustParse("100"), - }, - Requests: corev1.ResourceList{ - corev1.ResourceName(MLUResourceCount): resource.MustParse("1"), + corev1.ResourceName(MLUResourceCount): *resource.NewQuantity(1, resource.BinarySI), corev1.ResourceName(MLUResourceMemory): resource.MustParse("2048"), corev1.ResourceName(MLUResourceCores): resource.MustParse("100"), }, @@ -506,138 +493,14 @@ func setupTest(t *testing.T) (*corev1.Node, *corev1.Pod, func(), *fake.Clientset }, } - config := CambriconConfig{ - ResourceCountName: MLUResourceCount, - ResourceMemoryName: MLUResourceMemory, - ResourceCoreName: MLUResourceCores, - } - InitMLUDevice(config) - - _, err := clientset.CoreV1().Nodes().Create(ctx, node, metav1.CreateOptions{}) - if err != nil { - t.Fatalf("Failed to create node: %v", err) - } - - return node, pod, func() { - clientset.CoreV1().Nodes().Delete(ctx, node.Name, metav1.DeleteOptions{}) - }, clientset -} - -func Test_LockNode(t *testing.T) { - tests := []struct { - name string - annotations map[string]string - wantErr bool - }{ - { - name: "node is not locked", - annotations: map[string]string{}, - wantErr: false, - }, - { - name: "node is already locked within 2 minutes", - annotations: map[string]string{ - DsmluLockTime: time.Now().Add(-time.Minute).Format(time.RFC3339), - }, - wantErr: true, - }, - { - name: "lock time expired (more than 2 minutes)", - annotations: map[string]string{ - DsmluLockTime: time.Now().Add(-time.Hour).Format(time.RFC3339), - }, - wantErr: false, - }, - { - name: "invalid lock time format", - annotations: map[string]string{ - DsmluLockTime: "invalid-format", - }, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - node, pod, teardown, clientset := setupTest(t) - client.KubeClient = clientset - defer teardown() - - // Set up the node with the specified annotations. - node.Annotations = tt.annotations - - dev := InitMLUDevice(CambriconConfig{ - ResourceCountName: MLUResourceCount, - ResourceMemoryName: MLUResourceMemory, - ResourceCoreName: MLUResourceCores, - }) - - err := dev.LockNode(node, pod) - if (err != nil) != tt.wantErr { - t.Errorf("LockNode() error = %v, wantErr %v", err, tt.wantErr) - } - - // Optionally check if the node was correctly patched with the lock annotation. - if !tt.wantErr { - fetchedNode, _ := clientset.CoreV1().Nodes().Get(context.TODO(), node.Name, metav1.GetOptions{}) - if _, ok := fetchedNode.Annotations[DsmluLockTime]; !ok && !tt.wantErr { - t.Error("Expected node to be locked but it wasn't") - } - } - }) - } -} + dev := InitMLUDevice(config) + err = dev.ReleaseNodeLock(node, pod) + assert.NoError(t, err) -func Test_ReleaseNodeLock(t *testing.T) { - tests := []struct { - name string - args struct { - node corev1.Node - pod corev1.Pod - } - err error - }{ - { - name: "no annation", - args: struct { - node corev1.Node - pod corev1.Pod - }{ - node: corev1.Node{ - ObjectMeta: metav1.ObjectMeta{ - Name: "node-01", - }, - }, - pod: corev1.Pod{}, - }, - err: nil, - }, - { - name: "annation no lock value", - args: struct { - node corev1.Node - pod corev1.Pod - }{ - node: corev1.Node{ - ObjectMeta: metav1.ObjectMeta{ - Name: "node-02", - Annotations: map[string]string{ - "test": "test123", - }, - }, - }, - pod: corev1.Pod{}, - }, - err: nil, - }, - } - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - dev := CambriconDevices{} - result := dev.ReleaseNodeLock(&test.args.node, &test.args.pod) - assert.Equal(t, test.err, result) - }) - } + updated, err := client.KubeClient.CoreV1().Nodes().Get(context.Background(), "test-node", metav1.GetOptions{}) + assert.NoError(t, err) + _, ok := updated.Annotations[nodelock.NodeLockKey] + assert.False(t, ok) } func TestDevices_Fit(t *testing.T) {