Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cmd/scheduler/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
4 changes: 4 additions & 0 deletions pkg/scheduler/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
57 changes: 48 additions & 9 deletions pkg/scheduler/scheduler.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
}
}
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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):
}
}
Comment thread
archlitchi marked this conversation as resolved.
}

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
Expand Down Expand Up @@ -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)
Expand All @@ -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
}
Expand Down
130 changes: 130 additions & 0 deletions pkg/scheduler/scheduler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Comment thread
archlitchi marked this conversation as resolved.
config.NodeLockRetryTimeout = oldRetry
device.DevicesMap = oldDevicesMap
Comment thread
archlitchi marked this conversation as resolved.
close(s.stopCh)
Comment thread
archlitchi marked this conversation as resolved.
}
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")
}
11 changes: 10 additions & 1 deletion pkg/util/nodelock/nodelock.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ package nodelock

import (
"context"
"errors"
"fmt"
"os"
"strings"
Expand All @@ -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()
Expand Down Expand Up @@ -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) {
Expand Down
5 changes: 5 additions & 0 deletions pkg/util/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
8 changes: 8 additions & 0 deletions pkg/util/util.go
Original file line number Diff line number Diff line change
Expand Up @@ -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] != ""
}
Loading