Skip to content

Commit cea080f

Browse files
authored
fix(podgrouper): skip WorkloadRunner wrapper when selecting the grouping plugin (#2067)
Signed-off-by: gshaibi <gshaibi@nvidia.com>
1 parent e73b33d commit cea080f

6 files changed

Lines changed: 224 additions & 15 deletions

File tree

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
kind: Fixed
2+
body: |-
3+
Podgrouper skips WorkloadRunner wrapper so wrapped workloads keep their gang grouping

deployments/kai-scheduler/templates/rbac/podgrouper.yaml

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -253,6 +253,14 @@ rules:
253253
- create
254254
- patch
255255
- update
256+
- apiGroups:
257+
- nvidia.com
258+
resources:
259+
- dynamographdeployments
260+
verbs:
261+
- get
262+
- list
263+
- watch
256264
- apiGroups:
257265
- ray.io
258266
resources:
@@ -283,6 +291,7 @@ rules:
283291
- kartas
284292
- runaijobs
285293
- trainingworkloads
294+
- workloadrunners
286295
verbs:
287296
- get
288297
- list

pkg/podgrouper/podgrouper/hub/hub.go

Lines changed: 22 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ const (
4242
kindDistributedWorkload = "DistributedWorkload"
4343
kindInferenceWorkload = "InferenceWorkload"
4444
kindDistributedInferenceWorkload = "DistributedInferenceWorkload"
45+
kindWorkloadRunner = "WorkloadRunner"
4546
)
4647

4748
// +kubebuilder:rbac:groups=apps,resources=replicasets;statefulsets,verbs=get;list;watch
@@ -57,7 +58,9 @@ const (
5758
// +kubebuilder:rbac:groups=tekton.dev,resources=pipelineruns;taskruns,verbs=get;list;watch
5859
// +kubebuilder:rbac:groups=tekton.dev,resources=pipelineruns/finalizers;taskruns/finalizers,verbs=patch;update;create
5960
// +kubebuilder:rbac:groups=run.ai,resources=trainingworkloads;interactiveworkloads;distributedworkloads;inferenceworkloads;distributedinferenceworkloads,verbs=get;list;watch
61+
// +kubebuilder:rbac:groups=run.ai,resources=workloadrunners,verbs=get;list;watch
6062
// +kubebuilder:rbac:groups=run.ai,resources=kartas,verbs=get;list;watch
63+
// +kubebuilder:rbac:groups=nvidia.com,resources=dynamographdeployments,verbs=get;list;watch
6164
// +kubebuilder:rbac:groups=trainer.kubeflow.org,resources=trainjobs,verbs=get;list;watch
6265
// +kubebuilder:rbac:groups=trainer.kubeflow.org,resources=trainjobs/finalizers,verbs=patch;update;create
6366

@@ -298,7 +301,18 @@ func NewDefaultPluginsHub(kubeClient client.Client, searchForLegacyPodGroups,
298301
}: groveGrouper,
299302
}
300303

301-
skipTopOwnerGrouper := skiptopowner.NewSkipTopOwnerGrouper(kubeClient, defaultGrouper, table)
304+
hub := &DefaultPluginsHub{
305+
defaultGroupingHandler: defaultGrouper,
306+
customPlugins: table,
307+
}
308+
309+
// hub is captured by pointer and read only when the closure runs, so the skip
310+
// grouper still sees the table entries and Karta fallback assigned below.
311+
skipTopOwnerGrouper := skiptopowner.NewSkipTopOwnerGrouper(kubeClient, defaultGrouper,
312+
func(gvk metav1.GroupVersionKind) grouper.Grouper {
313+
return hub.GetPodGrouperPlugin(gvk)
314+
})
315+
302316
table[metav1.GroupVersionKind{
303317
Group: apiGroupArgo,
304318
Version: "v1alpha1",
@@ -337,10 +351,13 @@ func NewDefaultPluginsHub(kubeClient client.Client, searchForLegacyPodGroups,
337351
Kind: "DynamoGraphDeployment",
338352
}] = skipTopOwnerGrouper
339353

340-
hub := &DefaultPluginsHub{
341-
defaultGroupingHandler: defaultGrouper,
342-
customPlugins: table,
343-
}
354+
// WorkloadRunner is a kind-agnostic wrapper around an arbitrary workload template.
355+
// Skip it so the wrapped kind's plugin decides the grouping.
356+
table[metav1.GroupVersionKind{
357+
Group: apiGroupRunai,
358+
Version: "*",
359+
Kind: kindWorkloadRunner,
360+
}] = skipTopOwnerGrouper
344361

345362
if genericKartaFallback {
346363
hub.kartaFallbackPlugin = kartaplugin.NewKartaHub(kubeClient, defaultGrouper)

pkg/podgrouper/podgrouper/hub/hub_test.go

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,10 @@ import (
1010
. "github.com/onsi/gomega"
1111
kartav1alpha1 "github.com/run-ai/karta/pkg/api/runai/v1alpha1"
1212

13+
appsv1 "k8s.io/api/apps/v1"
14+
v1 "k8s.io/api/core/v1"
1315
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
16+
"k8s.io/apimachinery/pkg/apis/meta/v1/unstructured"
1417
"k8s.io/apimachinery/pkg/runtime"
1518
"k8s.io/apimachinery/pkg/types"
1619
"k8s.io/utils/ptr"
@@ -86,6 +89,28 @@ var _ = Describe("SupportedTypes", func() {
8689
Expect(hasPlugin).To(BeFalse())
8790
})
8891

92+
It("should return skipTopOwner plugin for WorkloadRunner", func() {
93+
gvk := metav1.GroupVersionKind{
94+
Group: "run.ai",
95+
Version: "v1alpha1",
96+
Kind: "WorkloadRunner",
97+
}
98+
plugin := hub.GetPodGrouperPlugin(gvk)
99+
Expect(plugin).NotTo(BeNil())
100+
Expect(plugin.Name()).To(BeEquivalentTo("SkipTopOwner Grouper"))
101+
})
102+
103+
It("should return skipTopOwner plugin for WorkloadRunner of any served version", func() {
104+
gvk := metav1.GroupVersionKind{
105+
Group: "run.ai",
106+
Version: "v2beta1",
107+
Kind: "WorkloadRunner",
108+
}
109+
plugin := hub.GetPodGrouperPlugin(gvk)
110+
Expect(plugin).NotTo(BeNil())
111+
Expect(plugin.Name()).To(BeEquivalentTo("SkipTopOwner Grouper"))
112+
})
113+
89114
It("should return skipTopOwner plugin for TrainJob", func() {
90115
gvk := metav1.GroupVersionKind{
91116
Group: "trainer.kubeflow.org",
@@ -155,6 +180,56 @@ var _ = Describe("SupportedTypes", func() {
155180
})
156181
})
157182

183+
Context("Skip Top Owner Resolution Tests", func() {
184+
It("should resolve the owner below a skipped WorkloadRunner through the hub", func() {
185+
statefulSet := &appsv1.StatefulSet{
186+
TypeMeta: metav1.TypeMeta{APIVersion: "apps/v1", Kind: "StatefulSet"},
187+
ObjectMeta: metav1.ObjectMeta{
188+
Name: "wrapped-sts",
189+
Namespace: "default",
190+
Labels: map[string]string{queueLabelKey: "wrapped-queue"},
191+
},
192+
}
193+
kubeClient := fake.NewFakeClient(statefulSet)
194+
hub := NewDefaultPluginsHub(
195+
kubeClient, false, false, false, queueLabelKey, nodePoolLabelKey, "", "",
196+
)
197+
198+
runner := &unstructured.Unstructured{}
199+
runner.SetAPIVersion("run.ai/v1alpha1")
200+
runner.SetKind("WorkloadRunner")
201+
runner.SetNamespace("default")
202+
runner.SetName("runner")
203+
204+
pod := &v1.Pod{
205+
TypeMeta: metav1.TypeMeta{APIVersion: "v1", Kind: "Pod"},
206+
ObjectMeta: metav1.ObjectMeta{Name: "test-pod", Namespace: "default"},
207+
}
208+
owners := []*metav1.PartialObjectMetadata{
209+
{
210+
TypeMeta: metav1.TypeMeta{APIVersion: "apps/v1", Kind: "StatefulSet"},
211+
ObjectMeta: metav1.ObjectMeta{Name: "wrapped-sts", Namespace: "default"},
212+
},
213+
{
214+
TypeMeta: metav1.TypeMeta{APIVersion: "run.ai/v1alpha1", Kind: "WorkloadRunner"},
215+
ObjectMeta: metav1.ObjectMeta{Name: "runner", Namespace: "default"},
216+
},
217+
}
218+
219+
plugin := hub.GetPodGrouperPlugin(metav1.GroupVersionKind{
220+
Group: apiGroupRunai, Version: "v1alpha1", Kind: "WorkloadRunner",
221+
})
222+
Expect(plugin.Name()).To(BeEquivalentTo("SkipTopOwner Grouper"))
223+
224+
metadata, err := plugin.GetPodGroupMetadata(runner, pod, owners...)
225+
226+
Expect(err).NotTo(HaveOccurred())
227+
Expect(metadata).NotTo(BeNil())
228+
Expect(metadata.Queue).To(Equal("wrapped-queue"))
229+
Expect(metadata.Owner.Kind).To(Equal("StatefulSet"))
230+
})
231+
})
232+
158233
Context("Wildcard Version Tests", func() {
159234
var (
160235
kubeClient client.Client

pkg/podgrouper/podgrouper/plugins/skiptopowner/skiptopowner.go

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -19,18 +19,21 @@ import (
1919
"github.com/kai-scheduler/KAI-scheduler/pkg/podgrouper/podgrouper/plugins/grouper"
2020
)
2121

22+
// GrouperResolver resolves the grouper handling a GVK, or nil when none matches.
23+
type GrouperResolver func(gvk metav1.GroupVersionKind) grouper.Grouper
24+
2225
type skipTopOwnerGrouper struct {
23-
client client.Client
24-
defaultPlugin *defaultgrouper.DefaultGrouper
25-
customPlugins map[metav1.GroupVersionKind]grouper.Grouper
26+
client client.Client
27+
defaultPlugin *defaultgrouper.DefaultGrouper
28+
resolveGrouper GrouperResolver
2629
}
2730

2831
func NewSkipTopOwnerGrouper(client client.Client, defaultGrouper *defaultgrouper.DefaultGrouper,
29-
customPlugins map[metav1.GroupVersionKind]grouper.Grouper) *skipTopOwnerGrouper {
32+
resolveGrouper GrouperResolver) *skipTopOwnerGrouper {
3033
return &skipTopOwnerGrouper{
31-
client: client,
32-
defaultPlugin: defaultGrouper,
33-
customPlugins: customPlugins,
34+
client: client,
35+
defaultPlugin: defaultGrouper,
36+
resolveGrouper: resolveGrouper,
3437
}
3538
}
3639

@@ -139,8 +142,10 @@ func (sk *skipTopOwnerGrouper) getSupportedTypePGMetadata(
139142
lastOwner *unstructured.Unstructured, pod *v1.Pod, otherOwners ...*metav1.PartialObjectMetadata,
140143
) (*podgroup.Metadata, error) {
141144
ownerKind := metav1.GroupVersionKind(lastOwner.GroupVersionKind())
142-
if grouper, found := sk.customPlugins[ownerKind]; found {
143-
return grouper.GetPodGroupMetadata(lastOwner, pod, otherOwners...)
145+
if sk.resolveGrouper != nil {
146+
if resolved := sk.resolveGrouper(ownerKind); resolved != nil {
147+
return resolved.GetPodGroupMetadata(lastOwner, pod, otherOwners...)
148+
}
144149
}
145150
return sk.defaultPlugin.GetPodGroupMetadata(lastOwner, pod, otherOwners...)
146151
}

pkg/podgrouper/podgrouper/plugins/skiptopowner/skiptopowner_test.go

Lines changed: 101 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ import (
1919
"sigs.k8s.io/controller-runtime/pkg/client"
2020
"sigs.k8s.io/controller-runtime/pkg/client/fake"
2121

22+
"github.com/kai-scheduler/KAI-scheduler/pkg/podgrouper/podgroup"
2223
"github.com/kai-scheduler/KAI-scheduler/pkg/podgrouper/podgrouper/plugins/constants"
2324
"github.com/kai-scheduler/KAI-scheduler/pkg/podgrouper/podgrouper/plugins/defaultgrouper"
2425
grouperplugin "github.com/kai-scheduler/KAI-scheduler/pkg/podgrouper/podgrouper/plugins/grouper"
@@ -65,7 +66,7 @@ var _ = Describe("SkipTopOwnerGrouper", func() {
6566
supportedTypes = map[metav1.GroupVersionKind]grouperplugin.Grouper{
6667
{Group: "", Version: "v1", Kind: "Pod"}: defaultGrouper,
6768
}
68-
plugin = NewSkipTopOwnerGrouper(client, defaultGrouper, supportedTypes)
69+
plugin = NewSkipTopOwnerGrouper(client, defaultGrouper, resolverFromMap(supportedTypes))
6970
})
7071

7172
Context("when last owner is a pod", func() {
@@ -430,6 +431,69 @@ var _ = Describe("SkipTopOwnerGrouper", func() {
430431
})
431432
})
432433

434+
Context("chained skip-top-owner delegation", func() {
435+
var (
436+
groveGrouper *recordingGrouper
437+
workloadRunner *unstructured.Unstructured
438+
pod *v1.Pod
439+
otherOwners []*metav1.PartialObjectMetadata
440+
)
441+
442+
BeforeEach(func() {
443+
groveGrouper = &recordingGrouper{name: "Grove Grouper"}
444+
supportedTypes[metav1.GroupVersionKind{
445+
Group: "grove.io", Version: "v1alpha1", Kind: "PodCliqueSet",
446+
}] = groveGrouper
447+
supportedTypes[metav1.GroupVersionKind{
448+
Group: "nvidia.com", Version: "*", Kind: "DynamoGraphDeployment",
449+
}] = plugin
450+
451+
// Dynamo 1.2.0+ serves the DGD as v1beta1, matched via the wildcard version.
452+
dgd := newUnstructured("nvidia.com/v1beta1", "DynamoGraphDeployment", "dgd")
453+
podCliqueSet := newUnstructured("grove.io/v1alpha1", "PodCliqueSet", "dgd-pcs")
454+
podCliqueSet.SetLabels(map[string]string{queueLabelKey: queueName})
455+
podClique := newUnstructured("grove.io/v1alpha1", "PodClique", "dgd-pcs-worker")
456+
workloadRunner = newUnstructured("run.ai/v1alpha1", "WorkloadRunner", "runner")
457+
458+
pod = examplePod.DeepCopy()
459+
pod.OwnerReferences = []metav1.OwnerReference{
460+
{Kind: "PodClique", APIVersion: "grove.io/v1alpha1", Name: podClique.GetName()},
461+
}
462+
463+
Expect(client.Create(context.TODO(), dgd)).To(Succeed())
464+
Expect(client.Create(context.TODO(), podCliqueSet)).To(Succeed())
465+
Expect(client.Create(context.TODO(), podClique)).To(Succeed())
466+
Expect(client.Create(context.TODO(), pod)).To(Succeed())
467+
pod.TypeMeta = metav1.TypeMeta{Kind: examplePod.Kind, APIVersion: examplePod.APIVersion}
468+
469+
otherOwners = []*metav1.PartialObjectMetadata{
470+
objectToPartial(podClique),
471+
objectToPartial(podCliqueSet),
472+
objectToPartial(dgd),
473+
objectToPartial(workloadRunner),
474+
}
475+
})
476+
477+
It("skips both WorkloadRunner and DGD and lands on the Grove grouper", func() {
478+
metadata, err := plugin.GetPodGroupMetadata(workloadRunner, pod, otherOwners...)
479+
480+
Expect(err).NotTo(HaveOccurred())
481+
Expect(metadata).NotTo(BeNil())
482+
Expect(metadata.Name).To(Equal("Grove Grouper"))
483+
Expect(groveGrouper.calledWith).NotTo(BeNil())
484+
Expect(groveGrouper.calledWith.GetKind()).To(Equal("PodCliqueSet"))
485+
})
486+
487+
It("propagates the skipped wrappers' labels down to the Grove owner", func() {
488+
workloadRunner.SetLabels(map[string]string{"runner-label": "runner-value"})
489+
490+
_, err := plugin.GetPodGroupMetadata(workloadRunner, pod, otherOwners...)
491+
492+
Expect(err).NotTo(HaveOccurred())
493+
Expect(groveGrouper.calledWith.GetLabels()).To(HaveKeyWithValue("runner-label", "runner-value"))
494+
})
495+
})
496+
433497
Context("default plugin usage", func() {
434498
It("uses default plugin when GVK does not have custom plugin", func() {
435499
// Use StatefulSet which is not in supportedTypes
@@ -492,6 +556,42 @@ var _ = Describe("SkipTopOwnerGrouper", func() {
492556

493557
})
494558

559+
type recordingGrouper struct {
560+
name string
561+
calledWith *unstructured.Unstructured
562+
}
563+
564+
func (r *recordingGrouper) Name() string { return r.name }
565+
566+
func (r *recordingGrouper) GetPodGroupMetadata(
567+
topOwner *unstructured.Unstructured, _ *v1.Pod, _ ...*metav1.PartialObjectMetadata,
568+
) (*podgroup.Metadata, error) {
569+
r.calledWith = topOwner
570+
return &podgroup.Metadata{Name: r.name}, nil
571+
}
572+
573+
func newUnstructured(apiVersion, kind, name string) *unstructured.Unstructured {
574+
obj := &unstructured.Unstructured{}
575+
obj.SetAPIVersion(apiVersion)
576+
obj.SetKind(kind)
577+
obj.SetName(name)
578+
obj.SetNamespace("default")
579+
return obj
580+
}
581+
582+
func resolverFromMap(table map[metav1.GroupVersionKind]grouperplugin.Grouper) GrouperResolver {
583+
return func(gvk metav1.GroupVersionKind) grouperplugin.Grouper {
584+
if g, found := table[gvk]; found {
585+
return g
586+
}
587+
gvk.Version = "*"
588+
if g, found := table[gvk]; found {
589+
return g
590+
}
591+
return nil
592+
}
593+
}
594+
495595
func objectToPartial(obj client.Object) *metav1.PartialObjectMetadata {
496596
objectMeta := metav1.ObjectMeta{
497597
Name: obj.GetName(),

0 commit comments

Comments
 (0)