diff --git a/cmd/scheduler/main.go b/cmd/scheduler/main.go index 5c977bb1a1..ef3e30c958 100644 --- a/cmd/scheduler/main.go +++ b/cmd/scheduler/main.go @@ -77,6 +77,7 @@ func init() { rootCmd.Flags().IntVar(&config.Timeout, "kube-timeout", client.DefaultTimeout, "Timeout to use while talking with kube-apiserver.") rootCmd.Flags().BoolVar(&enableProfiling, "profiling", false, "Enable pprof profiling via HTTP server") rootCmd.Flags().DurationVar(&config.NodeLockTimeout, "node-lock-timeout", time.Minute*5, "timeout for node locks") + rootCmd.Flags().DurationVar(&config.NodeLockRetryTimeout, "node-lock-retry-timeout", 28*time.Second, "timeout for retrying LockNode when contended by another PodGroup member (0 disables retry). Align the Extender's httpTimeout in KubeSchedulerConfiguration with this value.") rootCmd.Flags().BoolVar(&config.ForceOverwriteDefaultScheduler, "force-overwrite-default-scheduler", true, "Overwrite schedulerName in Pod Spec when set to the const DefaultSchedulerName in https://k8s.io/api/core/v1 package") rootCmd.Flags().BoolVar(&config.LeaderElect, "leader-elect", false, "The pod of hami-scheduler enable leader select") diff --git a/pkg/scheduler/config/config.go b/pkg/scheduler/config/config.go index c766e375e5..701a0c1bba 100644 --- a/pkg/scheduler/config/config.go +++ b/pkg/scheduler/config/config.go @@ -60,6 +60,10 @@ var ( // NodeLockTimeout is the timeout for node locks. NodeLockTimeout time.Duration + // NodeLockRetryTimeout is how long Bind retries LockNode when contended by + // another PodGroup member. Zero disables retry (fail-fast). + NodeLockRetryTimeout time.Duration + // If set to false, When Pod.Spec.SchedulerName equals to the const DefaultSchedulerName in k8s.io/api/core/v1 package, webhook will not overwrite it, default value is true. ForceOverwriteDefaultScheduler bool diff --git a/pkg/scheduler/scheduler.go b/pkg/scheduler/scheduler.go index 0008ccf1c1..8a87d5d8db 100644 --- a/pkg/scheduler/scheduler.go +++ b/pkg/scheduler/scheduler.go @@ -742,6 +742,50 @@ func (s *Scheduler) cleanupStalePodAllocation(pod *corev1.Pod) { } } +func (s *Scheduler) lockAllDevices(node *corev1.Node, pod *corev1.Pod) error { + for _, val := range device.GetDevices() { + if err := val.LockNode(node, pod); err != nil { + return err + } + } + return nil +} + +func (s *Scheduler) releaseAllDevices(node *corev1.Node, pod *corev1.Pod) { + for _, val := range device.GetDevices() { + if err := val.ReleaseNodeLock(node, pod); err != nil { + klog.ErrorS(err, "Failed to release node lock", "node", node.Name, "pod", klog.KObj(pod)) + } + } +} + +func (s *Scheduler) acquireNodeLocks(node *corev1.Node, pod *corev1.Pod) error { + if !util.IsPodGroupMember(pod) || config.NodeLockRetryTimeout <= 0 { + return s.lockAllDevices(node, pod) + } + + deadline := time.Now().Add(config.NodeLockRetryTimeout) + for { + err := s.lockAllDevices(node, pod) + if err == nil { + return nil + } + s.releaseAllDevices(node, pod) + if !nodelockutil.IsNodeLockContention(err) { + return err + } + if time.Now().After(deadline) { + return fmt.Errorf("timed out after %v waiting for node %s to be unlocked: %w", + config.NodeLockRetryTimeout, node.Name, nodelockutil.ErrNodeLockContention) + } + select { + case <-s.stopCh: + return fmt.Errorf("scheduler shutting down while waiting for node lock: %w", nodelockutil.ErrNodeLockContention) + case <-time.After(100 * time.Millisecond): + } + } +} + func (s *Scheduler) Bind(args extenderv1.ExtenderBindingArgs) (*extenderv1.ExtenderBindingResult, error) { klog.InfoS("Attempting to bind pod to node", "pod", args.PodName, "namespace", args.PodNamespace, "node", args.Node) var res *extenderv1.ExtenderBindingResult @@ -780,12 +824,9 @@ func (s *Scheduler) Bind(args extenderv1.ExtenderBindingArgs) (*extenderv1.Exten util.BindTimeAnnotations: strconv.FormatInt(time.Now().Unix(), 10), } - for _, val := range device.GetDevices() { - err = val.LockNode(node, current) - if err != nil { - klog.ErrorS(err, "Failed to lock node", "node", args.Node, "device", val) - goto ReleaseNodeLocks - } + if err = s.acquireNodeLocks(node, current); err != nil { + klog.ErrorS(err, "Failed to lock node", "node", args.Node, "pod", klog.KObj(current)) + goto ReleaseNodeLocks } err = util.PatchPodAnnotations(current, tmppatch) @@ -806,9 +847,7 @@ func (s *Scheduler) Bind(args extenderv1.ExtenderBindingArgs) (*extenderv1.Exten ReleaseNodeLocks: klog.InfoS("Release node locks", "node", args.Node) - for _, val := range device.GetDevices() { - val.ReleaseNodeLock(node, current) - } + s.releaseAllDevices(node, current) s.recordScheduleBindingResultEvent(current, EventReasonBindingFailed, []string{}, err) return &extenderv1.ExtenderBindingResult{Error: err.Error()}, nil } diff --git a/pkg/scheduler/scheduler_test.go b/pkg/scheduler/scheduler_test.go index 318db0d603..a7616ba3ca 100644 --- a/pkg/scheduler/scheduler_test.go +++ b/pkg/scheduler/scheduler_test.go @@ -1844,3 +1844,133 @@ func Test_Bind_DelPodOnGetNodeFailure(t *testing.T) { podsAfter, _ := s.podManager.ListPodsUID() require.Len(t, podsAfter, 0) } + +type bindLockMockDevice struct { + registerMockDevice + lockErr error + lockErrOnce bool + lockCalls atomic.Int32 + releaseCalls atomic.Int32 +} + +func (m *bindLockMockDevice) CommonWord() string { return "bind-lock-mock" } +func (m *bindLockMockDevice) LockNode(_ *corev1.Node, _ *corev1.Pod) error { + n := m.lockCalls.Add(1) + if m.lockErr != nil && (!m.lockErrOnce || n == 1) { + return m.lockErr + } + return nil +} +func (m *bindLockMockDevice) ReleaseNodeLock(_ *corev1.Node, _ *corev1.Pod) error { + m.releaseCalls.Add(1) + return nil +} + +var errContention = fmt.Errorf("contended: %w", nodelockutil.ErrNodeLockContention) + +func setupBindLockRetryTest(t *testing.T, retryTimeout time.Duration, pod *corev1.Pod, mock *bindLockMockDevice) (*Scheduler, extenderv1.ExtenderBindingArgs, func()) { + t.Helper() + + oldRetry := config.NodeLockRetryTimeout + config.NodeLockRetryTimeout = retryTimeout + oldDevicesMap := device.DevicesMap + device.DevicesMap = map[string]device.Devices{"bind-lock-mock": mock} + + s := NewScheduler() + cleanup := func() { + config.NodeLockRetryTimeout = oldRetry + device.DevicesMap = oldDevicesMap + close(s.stopCh) + } + scheme := runtime.NewScheme() + _ = corev1.AddToScheme(scheme) + s.eventRecorder = record.NewBroadcaster().NewRecorder(scheme, corev1.EventSource{}) + + node := &corev1.Node{ObjectMeta: metav1.ObjectMeta{Name: "node1"}} + fakeClient := fake.NewSimpleClientset(pod, node) + s.kubeClient = fakeClient + client.KubeClient = fakeClient + + informerFactory := informers.NewSharedInformerFactoryWithOptions(fakeClient, time.Hour) + require.NoError(t, informerFactory.Core().V1().Pods().Informer().GetIndexer().Add(pod)) + require.NoError(t, informerFactory.Core().V1().Nodes().Informer().GetIndexer().Add(node)) + s.podLister = informerFactory.Core().V1().Pods().Lister() + s.nodeLister = informerFactory.Core().V1().Nodes().Lister() + informerFactory.Start(s.stopCh) + informerFactory.WaitForCacheSync(s.stopCh) + + args := extenderv1.ExtenderBindingArgs{ + PodName: pod.Name, PodNamespace: pod.Namespace, PodUID: pod.UID, Node: "node1", + } + return s, args, cleanup +} + +func Test_Bind_NonPodGroupPodDoesNotRetry(t *testing.T) { + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "pod-nogroup", Namespace: "default", UID: types.UID("uid-nogroup"), + }, + } + mock := &bindLockMockDevice{lockErr: errContention, lockErrOnce: true} + s, args, cleanup := setupBindLockRetryTest(t, 5*time.Second, pod, mock) + defer cleanup() + + res, err := s.Bind(args) + require.NoError(t, err) + require.Contains(t, res.Error, "node lock contention") + require.Equal(t, int32(1), mock.lockCalls.Load(), + "non-PodGroup pod must not retry LockNode") +} + +func Test_Bind_PodGroupPodRetriesOnContention(t *testing.T) { + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "pod-gang", Namespace: "default", UID: types.UID("uid-gang"), + Labels: map[string]string{util.PodGroupLabel: "my-training-job"}, + }, + } + mock := &bindLockMockDevice{lockErr: errContention, lockErrOnce: true} + s, args, cleanup := setupBindLockRetryTest(t, 2*time.Second, pod, mock) + defer cleanup() + + s.Bind(args) + require.GreaterOrEqual(t, mock.lockCalls.Load(), int32(2), + "expected at least 2 LockNode calls (initial + retry)") +} + +func Test_Bind_PodGroupPodContendsUntilTimeout(t *testing.T) { + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "pod-timeout", Namespace: "default", UID: types.UID("uid-timeout"), + Labels: map[string]string{util.PodGroupLabel: "my-training-job"}, + }, + } + mock := &bindLockMockDevice{lockErr: errContention} + s, args, cleanup := setupBindLockRetryTest(t, 300*time.Millisecond, pod, mock) + defer cleanup() + + res, err := s.Bind(args) + require.NoError(t, err) + require.Contains(t, res.Error, "node lock contention", + "timeout error should wrap ErrNodeLockContention for observability") + require.GreaterOrEqual(t, mock.releaseCalls.Load(), int32(1), + "expected ReleaseNodeLock to be called at least once on timeout path") +} + +func Test_Bind_PodGroupPodNonContentionErrorDoesNotRetry(t *testing.T) { + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "pod-other-err", Namespace: "default", UID: types.UID("uid-other-err"), + Labels: map[string]string{util.PodGroupLabel: "my-training-job"}, + }, + } + mock := &bindLockMockDevice{lockErr: fmt.Errorf("apiserver 500"), lockErrOnce: true} + s, args, cleanup := setupBindLockRetryTest(t, 5*time.Second, pod, mock) + defer cleanup() + + res, err := s.Bind(args) + require.NoError(t, err) + require.Contains(t, res.Error, "apiserver 500") + require.Equal(t, int32(1), mock.lockCalls.Load(), + "non-contention error must not trigger retry") +} diff --git a/pkg/util/nodelock/nodelock.go b/pkg/util/nodelock/nodelock.go index 34871f2c86..0ba3998fff 100644 --- a/pkg/util/nodelock/nodelock.go +++ b/pkg/util/nodelock/nodelock.go @@ -18,6 +18,7 @@ package nodelock import ( "context" + "errors" "fmt" "os" "strings" @@ -40,6 +41,14 @@ const ( NodeLockSep = "," ) +// ErrNodeLockContention indicates the node lock is currently held by another +// valid pod. Callers may retry this error when the caller is a PodGroup member. +var ErrNodeLockContention = errors.New("node lock contention") + +func IsNodeLockContention(err error) bool { + return errors.Is(err, ErrNodeLockContention) +} + var ( // nodeLocks manages per-node locks for fine-grained concurrency control. nodeLocks = newNodeLockManager() @@ -246,7 +255,7 @@ func LockNode(nodeName string, lockname string, pods *corev1.Pod) error { return SetNodeLock(nodeName, lockname, pods) } - return fmt.Errorf("node %s has been locked within %v", nodeName, NodeLockTimeout) + return fmt.Errorf("node %s has been locked within %v: %w", nodeName, NodeLockTimeout, ErrNodeLockContention) } func ParseNodeLock(value string) (lockTime time.Time, ns, name string, err error) { diff --git a/pkg/util/types.go b/pkg/util/types.go index 780e4dfe1c..043ae18ea5 100644 --- a/pkg/util/types.go +++ b/pkg/util/types.go @@ -50,6 +50,11 @@ const ( HAMiComponentLabel = "app.kubernetes.io/component" // HAMiComponentScheduler the label value for hami-scheduler. HAMiComponentScheduler = "hami-scheduler" + + // PodGroupLabel is the label used by scheduler-plugins Coscheduling to mark + // a pod as a member of a PodGroup. See + // https://github.com/kubernetes-sigs/scheduler-plugins/blob/master/apis/scheduling/v1alpha1/types.go + PodGroupLabel = "scheduling.x-k8s.io/pod-group" ) var ( diff --git a/pkg/util/util.go b/pkg/util/util.go index ccf1b92423..b7ad6ab2f9 100644 --- a/pkg/util/util.go +++ b/pkg/util/util.go @@ -280,3 +280,11 @@ func IsPodTerminating(pod *corev1.Pod) bool { func AllContainersCreated(pod *corev1.Pod) bool { return len(pod.Status.ContainerStatuses) >= len(pod.Spec.Containers) } + +// Coscheduling PodGroup, based on the presence of the PodGroupLabel. +func IsPodGroupMember(pod *corev1.Pod) bool { + if pod == nil { + return false + } + return pod.Labels[PodGroupLabel] != "" +}