Skip to content

Commit 0622319

Browse files
committed
feat(binder): inherit fractional pod tolerations
Signed-off-by: Erez Freiberger <enoodle@gmail.com>
1 parent b8730e3 commit 0622319

3 files changed

Lines changed: 61 additions & 16 deletions

File tree

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
kind: Added
2+
body: |-
3+
Reservation pods inherit fractional pod tolerations

pkg/binder/binding/resourcereservation/resource_reservation.go

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -252,7 +252,7 @@ func (rsc *service) ReserveGpuDevice(ctx context.Context, pod *v1.Pod, nodeName
252252
rsc.gpuGroupMutex.LockMutexForGroup(gpuGroup)
253253
defer rsc.gpuGroupMutex.ReleaseMutex(gpuGroup)
254254

255-
gpuIndex, err := rsc.acquireGPUIndexByGroup(ctx, nodeName, gpuGroup)
255+
gpuIndex, err := rsc.acquireGPUIndexByGroup(ctx, pod, nodeName, gpuGroup)
256256
if err != nil {
257257
return unknownGpuIndicator, err
258258
}
@@ -332,15 +332,17 @@ func escapeJSONPointer(s string) string {
332332
return s
333333
}
334334

335-
func (rsc *service) acquireGPUIndexByGroup(ctx context.Context, nodeName, gpuGroup string) (string, error) {
335+
func (rsc *service) acquireGPUIndexByGroup(
336+
ctx context.Context, sourcePod *v1.Pod, nodeName, gpuGroup string,
337+
) (string, error) {
336338
gpuIndex, err := rsc.findGPUIndexByGroup(gpuGroup)
337339
if err != nil {
338340
return "", err
339341
}
340342
if gpuIndex != "" {
341343
return gpuIndex, err
342344
}
343-
return rsc.createGPUReservationPodAndGetIndex(ctx, nodeName, gpuGroup)
345+
return rsc.createGPUReservationPodAndGetIndex(ctx, sourcePod, nodeName, gpuGroup)
344346
}
345347

346348
func (rsc *service) findGPUIndexByGroup(gpuGroup string) (
@@ -366,10 +368,12 @@ func (rsc *service) findGPUIndexByGroup(gpuGroup string) (
366368
return gpuIndex, nil
367369
}
368370

369-
func (rsc *service) createGPUReservationPodAndGetIndex(ctx context.Context, nodeName, gpuGroup string) (
371+
func (rsc *service) createGPUReservationPodAndGetIndex(
372+
ctx context.Context, sourcePod *v1.Pod, nodeName, gpuGroup string,
373+
) (
370374
gpuIndex string, err error) {
371375
logger := log.FromContext(ctx)
372-
pod, err := rsc.createGPUReservationPod(ctx, nodeName, gpuGroup)
376+
pod, err := rsc.createGPUReservationPod(ctx, sourcePod, nodeName, gpuGroup)
373377
if err != nil {
374378
return unknownGpuIndicator, err
375379
}
@@ -421,7 +425,9 @@ func (rsc *service) deleteReservationPod(ctx context.Context, pod *v1.Pod) error
421425
return nil
422426
}
423427

424-
func (rsc *service) createGPUReservationPod(ctx context.Context, nodeName, gpuGroup string) (*v1.Pod, error) {
428+
func (rsc *service) createGPUReservationPod(
429+
ctx context.Context, sourcePod *v1.Pod, nodeName, gpuGroup string,
430+
) (*v1.Pod, error) {
425431
logger := log.FromContext(ctx)
426432
if rsc.isScalingUp(ctx) {
427433
return nil, fmt.Errorf("cluster is scaling up, could not create reservation pod")
@@ -450,7 +456,7 @@ func (rsc *service) createGPUReservationPod(ctx context.Context, nodeName, gpuGr
450456
}
451457
}
452458

453-
pod, err := rsc.createResourceReservationPod(nodeName, gpuGroup, podName, resources)
459+
pod, err := rsc.createResourceReservationPod(sourcePod, nodeName, gpuGroup, podName, resources)
454460
if err != nil {
455461
// The reservation pod name is deterministic per (node, gpu-group). AlreadyExists
456462
// means another actor (a concurrent bind, a retry, or another binder replica)
@@ -525,8 +531,14 @@ func (rsc *service) waitForGPUReservationPodAllocation(
525531
}
526532

527533
func (rsc *service) createResourceReservationPod(
528-
nodeName, gpuGroup, podName string, resources v1.ResourceRequirements,
534+
sourcePod *v1.Pod, nodeName, gpuGroup, podName string, resources v1.ResourceRequirements,
529535
) (*v1.Pod, error) {
536+
var tolerations []v1.Toleration
537+
if sourcePod != nil {
538+
sourcePodSpec := sourcePod.Spec.DeepCopy()
539+
tolerations = sourcePodSpec.Tolerations
540+
}
541+
530542
podSpec := &v1.Pod{
531543
ObjectMeta: metav1.ObjectMeta{
532544
Name: podName,
@@ -540,7 +552,8 @@ func (rsc *service) createResourceReservationPod(
540552
},
541553
},
542554
Spec: v1.PodSpec{
543-
NodeName: nodeName,
555+
NodeName: nodeName,
556+
Tolerations: tolerations,
544557
RuntimeClassName: func() *string {
545558
if len(rsc.runtimeClassName) == 0 {
546559
return nil

pkg/binder/binding/resourcereservation/resource_reservation_test.go

Lines changed: 36 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -910,7 +910,7 @@ var _ = Describe("ResourceReservationService", func() {
910910
gpuGroup := "test-group"
911911
nodeName := "node-test"
912912

913-
pod, err := rsc.createResourceReservationPod(nodeName, gpuGroup, podName, resources)
913+
pod, err := rsc.createResourceReservationPod(nil, nodeName, gpuGroup, podName, resources)
914914
Expect(err).To(BeNil())
915915
Expect(pod).NotTo(BeNil())
916916

@@ -953,6 +953,35 @@ var _ = Describe("ResourceReservationService", func() {
953953
Expect(container.Env).To(ContainElement(Equal(podNameEnv)))
954954
Expect(container.Env).To(ContainElement(Equal(podNamespaceEnv)))
955955
})
956+
957+
It("should copy tolerations from the source pod", func() {
958+
rsc := &service{
959+
namespace: "kai-resource-reservation",
960+
appLabelValue: "kai-reservation",
961+
serviceAccountName: "kai-sa",
962+
reservationPodImage: "nvidia/kai-reservation:latest",
963+
kubeClient: fake.NewClientBuilder().WithScheme(testScheme).Build(),
964+
}
965+
sourcePod := &v1.Pod{
966+
Spec: v1.PodSpec{
967+
Tolerations: []v1.Toleration{{
968+
Key: "hpc",
969+
Operator: v1.TolerationOpEqual,
970+
Value: "true",
971+
Effect: v1.TaintEffectNoExecute,
972+
}},
973+
},
974+
}
975+
976+
pod, err := rsc.createResourceReservationPod(
977+
sourcePod, "node-test", "test-group", "reservation-test", v1.ResourceRequirements{},
978+
)
979+
Expect(err).To(Succeed())
980+
Expect(pod.Spec.Tolerations).To(Equal(sourcePod.Spec.Tolerations))
981+
982+
sourcePod.Spec.Tolerations[0].Value = "changed"
983+
Expect(pod.Spec.Tolerations[0].Value).To(Equal("true"))
984+
})
956985
})
957986

958987
Context("RemovePodGpuGroupsConnection", func() {
@@ -1116,7 +1145,7 @@ var _ = Describe("ResourceReservationService", func() {
11161145
scalingPodNamespace: scalingPodsNamespace,
11171146
}
11181147

1119-
pod, err := rsc.createGPUReservationPod(context.TODO(), "test-node", "test-gpu-group")
1148+
pod, err := rsc.createGPUReservationPod(context.TODO(), nil, "test-node", "test-gpu-group")
11201149
Expect(err).To(BeNil())
11211150
Expect(pod).NotTo(BeNil())
11221151

@@ -1146,7 +1175,7 @@ var _ = Describe("ResourceReservationService", func() {
11461175
scalingPodNamespace: scalingPodsNamespace,
11471176
}
11481177

1149-
pod, err := rsc.createGPUReservationPod(context.TODO(), "test-node", "test-gpu-group")
1178+
pod, err := rsc.createGPUReservationPod(context.TODO(), nil, "test-node", "test-gpu-group")
11501179
Expect(err).To(BeNil())
11511180
Expect(pod).NotTo(BeNil())
11521181

@@ -1188,7 +1217,7 @@ var _ = Describe("ResourceReservationService", func() {
11881217
scalingPodNamespace: scalingPodsNamespace,
11891218
}
11901219

1191-
pod, err := rsc.createGPUReservationPod(context.TODO(), "test-node", "test-gpu-group")
1220+
pod, err := rsc.createGPUReservationPod(context.TODO(), nil, "test-node", "test-gpu-group")
11921221
Expect(err).To(BeNil())
11931222
Expect(pod).NotTo(BeNil())
11941223

@@ -1227,7 +1256,7 @@ var _ = Describe("ResourceReservationService", func() {
12271256
scalingPodNamespace: scalingPodsNamespace,
12281257
}
12291258

1230-
pod, err := rsc.createGPUReservationPod(context.TODO(), "test-node", "test-gpu-group")
1259+
pod, err := rsc.createGPUReservationPod(context.TODO(), nil, "test-node", "test-gpu-group")
12311260
Expect(err).To(BeNil())
12321261
Expect(pod).NotTo(BeNil())
12331262

@@ -1505,7 +1534,7 @@ var _ = Describe("Race condition: reservation pod deleted during concurrent bind
15051534
nil, podSecCtx, containerSecCtx)
15061535

15071536
pod, err := svc.createResourceReservationPod(
1508-
nodeName, gpuGroup, "test-reservation-pod",
1537+
nil, nodeName, gpuGroup, "test-reservation-pod",
15091538
v1.ResourceRequirements{
15101539
Limits: v1.ResourceList{constants.NvidiaGpuResource: *resource.NewQuantity(1, resource.DecimalSI)},
15111540
Requests: v1.ResourceList{constants.NvidiaGpuResource: *resource.NewQuantity(1, resource.DecimalSI)},
@@ -1525,7 +1554,7 @@ var _ = Describe("Race condition: reservation pod deleted during concurrent bind
15251554
nil, nil, nil)
15261555

15271556
pod, err := svc.createResourceReservationPod(
1528-
nodeName, gpuGroup, "test-reservation-pod",
1557+
nil, nodeName, gpuGroup, "test-reservation-pod",
15291558
v1.ResourceRequirements{
15301559
Limits: v1.ResourceList{constants.NvidiaGpuResource: *resource.NewQuantity(1, resource.DecimalSI)},
15311560
Requests: v1.ResourceList{constants.NvidiaGpuResource: *resource.NewQuantity(1, resource.DecimalSI)},

0 commit comments

Comments
 (0)