Skip to content
This repository was archived by the owner on Oct 9, 2023. It is now read-only.

Commit f5f4182

Browse files
ByronHsubyhsu
andauthored
Override primary container name instead of flyte generated name (#340)
Signed-off-by: byhsu <byhsu@linkedin.com> Co-authored-by: byhsu <byhsu@linkedin.com>
1 parent 435436b commit f5f4182

6 files changed

Lines changed: 21 additions & 23 deletions

File tree

go/tasks/pluginmachinery/flytek8s/pod_helper.go

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -255,20 +255,20 @@ func ApplyFlytePodConfiguration(ctx context.Context, tCtx pluginsCore.TaskExecut
255255

256256
// ToK8sPodSpec builds a PodSpec and ObjectMeta based on the definition passed by the TaskExecutionContext. This
257257
// involves parsing the raw PodSpec definition and applying all Flyte configuration options.
258-
func ToK8sPodSpec(ctx context.Context, tCtx pluginsCore.TaskExecutionContext) (*v1.PodSpec, *metav1.ObjectMeta, error) {
258+
func ToK8sPodSpec(ctx context.Context, tCtx pluginsCore.TaskExecutionContext) (*v1.PodSpec, *metav1.ObjectMeta, string, error) {
259259
// build raw PodSpec and ObjectMeta
260260
podSpec, objectMeta, primaryContainerName, err := BuildRawPod(ctx, tCtx)
261261
if err != nil {
262-
return nil, nil, err
262+
return nil, nil, "", err
263263
}
264264

265265
// add flyte configuration
266266
podSpec, objectMeta, err = ApplyFlytePodConfiguration(ctx, tCtx, podSpec, objectMeta, primaryContainerName)
267267
if err != nil {
268-
return nil, nil, err
268+
return nil, nil, "", err
269269
}
270270

271-
return podSpec, objectMeta, nil
271+
return podSpec, objectMeta, primaryContainerName, nil
272272
}
273273

274274
// getBasePodTemplate attempts to retrieve the PodTemplate to use as the base for k8s Pod configuration. This value can

go/tasks/pluginmachinery/flytek8s/pod_helper_test.go

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -324,7 +324,7 @@ func toK8sPodInterruptible(t *testing.T) {
324324
},
325325
})
326326

327-
p, _, err := ToK8sPodSpec(ctx, x)
327+
p, _, _, err := ToK8sPodSpec(ctx, x)
328328
assert.NoError(t, err)
329329
assert.Len(t, p.Tolerations, 2)
330330
assert.Equal(t, "x/flyte", p.Tolerations[1].Key)
@@ -391,7 +391,7 @@ func TestToK8sPod(t *testing.T) {
391391
},
392392
})
393393

394-
p, _, err := ToK8sPodSpec(ctx, x)
394+
p, _, _, err := ToK8sPodSpec(ctx, x)
395395
assert.NoError(t, err)
396396
assert.Equal(t, len(p.Tolerations), 1)
397397
})
@@ -408,7 +408,7 @@ func TestToK8sPod(t *testing.T) {
408408
},
409409
})
410410

411-
p, _, err := ToK8sPodSpec(ctx, x)
411+
p, _, _, err := ToK8sPodSpec(ctx, x)
412412
assert.NoError(t, err)
413413
assert.Equal(t, len(p.Tolerations), 0)
414414
assert.Equal(t, "some-acceptable-name", p.Containers[0].Name)
@@ -435,7 +435,7 @@ func TestToK8sPod(t *testing.T) {
435435
DefaultMemoryRequest: resource.MustParse("1024Mi"),
436436
}))
437437

438-
p, _, err := ToK8sPodSpec(ctx, x)
438+
p, _, _, err := ToK8sPodSpec(ctx, x)
439439
assert.NoError(t, err)
440440
assert.Equal(t, 1, len(p.NodeSelector))
441441
assert.Equal(t, "myScheduler", p.SchedulerName)
@@ -452,7 +452,7 @@ func TestToK8sPod(t *testing.T) {
452452
}))
453453

454454
x := dummyExecContext(&v1.ResourceRequirements{})
455-
p, _, err := ToK8sPodSpec(ctx, x)
455+
p, _, _, err := ToK8sPodSpec(ctx, x)
456456
assert.NoError(t, err)
457457
assert.NotNil(t, p.SecurityContext)
458458
assert.Equal(t, *p.SecurityContext.RunAsGroup, v)
@@ -464,7 +464,7 @@ func TestToK8sPod(t *testing.T) {
464464
EnableHostNetworkingPod: &enabled,
465465
}))
466466
x := dummyExecContext(&v1.ResourceRequirements{})
467-
p, _, err := ToK8sPodSpec(ctx, x)
467+
p, _, _, err := ToK8sPodSpec(ctx, x)
468468
assert.NoError(t, err)
469469
assert.True(t, p.HostNetwork)
470470
})
@@ -475,15 +475,15 @@ func TestToK8sPod(t *testing.T) {
475475
EnableHostNetworkingPod: &enabled,
476476
}))
477477
x := dummyExecContext(&v1.ResourceRequirements{})
478-
p, _, err := ToK8sPodSpec(ctx, x)
478+
p, _, _, err := ToK8sPodSpec(ctx, x)
479479
assert.NoError(t, err)
480480
assert.False(t, p.HostNetwork)
481481
})
482482

483483
t.Run("skipSettingHostNetwork", func(t *testing.T) {
484484
assert.NoError(t, config.SetK8sPluginConfig(&config.K8sPluginConfig{}))
485485
x := dummyExecContext(&v1.ResourceRequirements{})
486-
p, _, err := ToK8sPodSpec(ctx, x)
486+
p, _, _, err := ToK8sPodSpec(ctx, x)
487487
assert.NoError(t, err)
488488
assert.False(t, p.HostNetwork)
489489
})
@@ -517,7 +517,7 @@ func TestToK8sPod(t *testing.T) {
517517
}))
518518

519519
x := dummyExecContext(&v1.ResourceRequirements{})
520-
p, _, err := ToK8sPodSpec(ctx, x)
520+
p, _, _, err := ToK8sPodSpec(ctx, x)
521521
assert.NoError(t, err)
522522
assert.NotNil(t, p.DNSConfig)
523523
assert.Equal(t, []string{"8.8.8.8", "8.8.4.4"}, p.DNSConfig.Nameservers)

go/tasks/plugins/k8s/kfoperators/common/common_operator.go

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -166,17 +166,15 @@ func GetLogs(taskType string, name string, namespace string,
166166
return taskLogs, nil
167167
}
168168

169-
func OverrideDefaultContainerName(taskCtx pluginsCore.TaskExecutionContext, podSpec *v1.PodSpec,
170-
defaultContainerName string) {
169+
func OverridePrimaryContainerName(podSpec *v1.PodSpec, primaryContainerName string, defaultContainerName string) {
171170
// Pytorch operator forces pod to have container named 'pytorch'
172171
// https://github.com/kubeflow/pytorch-operator/blob/037cd1b18eb77f657f2a4bc8a8334f2a06324b57/pkg/apis/pytorch/validation/validation.go#L54-L62
173172
// Tensorflow operator forces pod to have container named 'tensorflow'
174173
// https://github.com/kubeflow/tf-operator/blob/984adc287e6fe82841e4ca282dc9a2cbb71e2d4a/pkg/apis/tensorflow/validation/validation.go#L55-L63
175174
// hence we have to override the name set here
176175
// https://github.com/flyteorg/flyteplugins/blob/209c52d002b4e6a39be5d175bc1046b7e631c153/go/tasks/pluginmachinery/flytek8s/container_helper.go#L116
177-
flyteDefaultContainerName := taskCtx.TaskExecutionMetadata().GetTaskExecutionID().GetGeneratedName()
178176
for idx, c := range podSpec.Containers {
179-
if c.Name == flyteDefaultContainerName {
177+
if c.Name == primaryContainerName {
180178
podSpec.Containers[idx].Name = defaultContainerName
181179
return
182180
}

go/tasks/plugins/k8s/kfoperators/mpi/mpi.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -62,11 +62,11 @@ func (mpiOperatorResourceHandler) BuildResource(ctx context.Context, taskCtx plu
6262
launcherReplicas := mpiTaskExtraArgs.GetNumLauncherReplicas()
6363
slots := mpiTaskExtraArgs.GetSlots()
6464

65-
podSpec, objectMeta, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
65+
podSpec, objectMeta, primaryContainerName, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
6666
if err != nil {
6767
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "Unable to create pod spec: [%v]", err.Error())
6868
}
69-
common.OverrideDefaultContainerName(taskCtx, podSpec, kubeflowv1.MPIJobDefaultContainerName)
69+
common.OverridePrimaryContainerName(podSpec, primaryContainerName, kubeflowv1.MPIJobDefaultContainerName)
7070

7171
// workersPodSpec is deepCopy of podSpec submitted by flyte
7272
// WorkerPodSpec doesn't need any Argument & command. It will be trigger from launcher pod

go/tasks/plugins/k8s/kfoperators/pytorch/pytorch.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -62,11 +62,11 @@ func (pytorchOperatorResourceHandler) BuildResource(ctx context.Context, taskCtx
6262
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "invalid TaskSpecification [%v], Err: [%v]", taskTemplate.GetCustom(), err.Error())
6363
}
6464

65-
podSpec, objectMeta, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
65+
podSpec, objectMeta, primaryContainerName, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
6666
if err != nil {
6767
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "Unable to create pod spec: [%v]", err.Error())
6868
}
69-
common.OverrideDefaultContainerName(taskCtx, podSpec, kubeflowv1.PytorchJobDefaultContainerName)
69+
common.OverridePrimaryContainerName(podSpec, primaryContainerName, kubeflowv1.PytorchJobDefaultContainerName)
7070

7171
workers := pytorchTaskExtraArgs.GetWorkers()
7272
if workers == 0 {

go/tasks/plugins/k8s/kfoperators/tensorflow/tensorflow.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -62,11 +62,11 @@ func (tensorflowOperatorResourceHandler) BuildResource(ctx context.Context, task
6262
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "invalid TaskSpecification [%v], Err: [%v]", taskTemplate.GetCustom(), err.Error())
6363
}
6464

65-
podSpec, objectMeta, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
65+
podSpec, objectMeta, primaryContainerName, err := flytek8s.ToK8sPodSpec(ctx, taskCtx)
6666
if err != nil {
6767
return nil, flyteerr.Errorf(flyteerr.BadTaskSpecification, "Unable to create pod spec: [%v]", err.Error())
6868
}
69-
common.OverrideDefaultContainerName(taskCtx, podSpec, kubeflowv1.TFJobDefaultContainerName)
69+
common.OverridePrimaryContainerName(podSpec, primaryContainerName, kubeflowv1.TFJobDefaultContainerName)
7070

7171
workers := tensorflowTaskExtraArgs.GetWorkers()
7272
psReplicas := tensorflowTaskExtraArgs.GetPsReplicas()

0 commit comments

Comments
 (0)