From 41de0627a896b6bb3de29a9cce2f1100e23d1ba3 Mon Sep 17 00:00:00 2001 From: Gaurav-205 Date: Tue, 28 Jul 2026 15:10:04 +0530 Subject: [PATCH] fix: process init containers in admission mutation and quota checks Signed-off-by: Gaurav-205 --- pkg/scheduler/webhook.go | 22 ++++++++ pkg/scheduler/webhook_test.go | 101 ++++++++++++++++++++++++++++++++++ 2 files changed, 123 insertions(+) diff --git a/pkg/scheduler/webhook.go b/pkg/scheduler/webhook.go index 3f749e8f46..12000f980d 100644 --- a/pkg/scheduler/webhook.go +++ b/pkg/scheduler/webhook.go @@ -71,6 +71,17 @@ func (h *webhook) Handle(_ context.Context, req admission.Request) admission.Res klog.V(5).Infof(template, pod.Namespace, pod.Name, pod.UID) privilegedName, hasPrivileged := privilegedContainerName(pod) hasResource := false + for idx := range pod.Spec.InitContainers { + c := &pod.Spec.InitContainers[idx] + for _, val := range device.GetDevices() { + found, err := val.MutateAdmission(c, pod) + if err != nil { + klog.Errorf("validating pod failed:%s", err.Error()) + return admission.Errored(http.StatusInternalServerError, err) + } + hasResource = hasResource || found + } + } for idx := range pod.Spec.Containers { c := &pod.Spec.Containers[idx] for _, val := range device.GetDevices() { @@ -153,6 +164,17 @@ func fitResourceQuota(pod *corev1.Pod) bool { } return 0, false } + for _, ctr := range pod.Spec.InitContainers { + req, ok := getRequest(&ctr, resourceName) + if ok { + if memReq, ok := getRequest(&ctr, memResourceName); ok { + memoryReq += memReq * req + } + if coreReq, ok := getRequest(&ctr, coreResourceName); ok { + coresReq += coreReq * req + } + } + } for _, ctr := range pod.Spec.Containers { req, ok := getRequest(&ctr, resourceName) if ok { diff --git a/pkg/scheduler/webhook_test.go b/pkg/scheduler/webhook_test.go index a03e488310..1369c5e113 100644 --- a/pkg/scheduler/webhook_test.go +++ b/pkg/scheduler/webhook_test.go @@ -811,3 +811,104 @@ func TestPrivilegedContainerDenied(t *testing.T) { }) } } + +func TestInitContainerOnlyGPUAdmission(t *testing.T) { + prevSchedulerName := config.SchedulerName + prevForceOverwrite := config.ForceOverwriteDefaultScheduler + t.Cleanup(func() { + config.SchedulerName = prevSchedulerName + config.ForceOverwriteDefaultScheduler = prevForceOverwrite + }) + + config.SchedulerName = "hami-scheduler" + config.ForceOverwriteDefaultScheduler = false + + sConfig := &config.Config{ + NvidiaConfig: nvidia.NvidiaConfig{ + ResourceCountName: "hami.io/gpu", + ResourceMemoryName: "hami.io/gpumem", + ResourceMemoryPercentageName: "hami.io/gpumem-percentage", + ResourceCoreName: "hami.io/gpucores", + DefaultMemory: 0, + DefaultCores: 0, + DefaultGPUNum: 1, + }, + } + + if err := config.InitDevicesWithConfig(sConfig); err != nil { + t.Fatalf("Failed to initialize devices with config: %v", err) + } + + // Pod with GPU resource in init container only + pod := &corev1.Pod{ + ObjectMeta: metav1.ObjectMeta{ + Name: "init-gpu-pod", + Namespace: "default", + }, + Spec: corev1.PodSpec{ + InitContainers: []corev1.Container{ + { + Name: "gpu-init", + Image: "cuda-init", + Resources: corev1.ResourceRequirements{ + Limits: corev1.ResourceList{ + "hami.io/gpu": resource.MustParse("1"), + }, + }, + }, + }, + Containers: []corev1.Container{ + { + Name: "cpu-app", + Image: "busybox", + Resources: corev1.ResourceRequirements{ + Limits: corev1.ResourceList{ + corev1.ResourceCPU: resource.MustParse("1"), + }, + }, + }, + }, + }, + } + + scheme := runtime.NewScheme() + corev1.AddToScheme(scheme) + codec := serializer.NewCodecFactory(scheme).LegacyCodec(corev1.SchemeGroupVersion) + podBytes, err := runtime.Encode(codec, pod) + if err != nil { + t.Fatalf("Error encoding pod: %v", err) + } + + req := admission.Request{ + AdmissionRequest: admissionv1.AdmissionRequest{ + UID: "init-gpu-uid", + Namespace: "default", + Name: "init-gpu-pod", + Object: runtime.RawExtension{ + Raw: podBytes, + }, + }, + } + + wh, err := NewWebHook() + if err != nil { + t.Fatalf("Error creating WebHook: %v", err) + } + + resp := wh.Handle(context.Background(), req) + if !resp.Allowed { + t.Fatalf("Expected init GPU pod to be allowed, but got denied: %+v", resp.Result) + } + + // Verify that schedulerName was patched to hami-scheduler for init container GPU request + found := false + for _, patch := range resp.Patches { + if patch.Path == "/spec/schedulerName" && patch.Value == config.SchedulerName { + found = true + break + } + } + if !found { + t.Fatalf("Expected schedulerName patch to %q for init container GPU request, got patches: %+v", config.SchedulerName, resp.Patches) + } +}