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

Commit cc2b8c6

Browse files
authored
Supporting interruptible for map tasks (#253)
* implemented IsInterruptible for SubTaskExecutionMetadata Signed-off-by: Daniel Rammer <daniel@union.ai> * fixed possible race condition Signed-off-by: Daniel Rammer <daniel@union.ai> * fixed unit tests Signed-off-by: Daniel Rammer <daniel@union.ai> * fixed lint issue Signed-off-by: Daniel Rammer <daniel@union.ai> * updated TODO documentation Signed-off-by: Daniel Rammer <daniel@union.ai> * changed context on NewCompactArray error log Signed-off-by: Daniel Rammer <daniel@union.ai> * fixing retry attempt calculation on abort Signed-off-by: Daniel Rammer <daniel@union.ai>
1 parent 6f534eb commit cc2b8c6

10 files changed

Lines changed: 95 additions & 62 deletions

File tree

go/tasks/pluginmachinery/core/exec_metadata.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,4 +44,5 @@ type TaskExecutionMetadata interface {
4444
GetSecurityContext() core.SecurityContext
4545
IsInterruptible() bool
4646
GetPlatformResources() *v1.ResourceRequirements
47+
GetInterruptibleFailureThreshold() uint32
4748
}

go/tasks/pluginmachinery/core/mocks/task_execution_metadata.go

Lines changed: 32 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

go/tasks/pluginmachinery/flytek8s/pod_helper.go

Lines changed: 4 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -63,18 +63,11 @@ func ApplyInterruptibleNodeAffinity(interruptible bool, podSpec *v1.PodSpec) {
6363
// UpdatePod updates the base pod spec used to execute tasks. This is configured with plugins and task metadata-specific options
6464
func UpdatePod(taskExecutionMetadata pluginsCore.TaskExecutionMetadata,
6565
resourceRequirements []v1.ResourceRequirements, podSpec *v1.PodSpec) {
66-
UpdatePodWithInterruptibleFlag(taskExecutionMetadata, resourceRequirements, podSpec, false)
67-
}
68-
69-
// UpdatePodWithInterruptibleFlag updates the base pod spec used to execute tasks. This is configured with plugins and task metadata-specific options
70-
func UpdatePodWithInterruptibleFlag(taskExecutionMetadata pluginsCore.TaskExecutionMetadata,
71-
resourceRequirements []v1.ResourceRequirements, podSpec *v1.PodSpec, omitInterruptible bool) {
72-
isInterruptible := !omitInterruptible && taskExecutionMetadata.IsInterruptible()
7366
if len(podSpec.RestartPolicy) == 0 {
7467
podSpec.RestartPolicy = v1.RestartPolicyNever
7568
}
7669
podSpec.Tolerations = append(
77-
GetPodTolerations(isInterruptible, resourceRequirements...), podSpec.Tolerations...)
70+
GetPodTolerations(taskExecutionMetadata.IsInterruptible(), resourceRequirements...), podSpec.Tolerations...)
7871

7972
if len(podSpec.ServiceAccountName) == 0 {
8073
podSpec.ServiceAccountName = taskExecutionMetadata.GetK8sServiceAccount()
@@ -83,7 +76,7 @@ func UpdatePodWithInterruptibleFlag(taskExecutionMetadata pluginsCore.TaskExecut
8376
podSpec.SchedulerName = config.GetK8sPluginConfig().SchedulerName
8477
}
8578
podSpec.NodeSelector = utils.UnionMaps(podSpec.NodeSelector, config.GetK8sPluginConfig().DefaultNodeSelector)
86-
if isInterruptible {
79+
if taskExecutionMetadata.IsInterruptible() {
8780
podSpec.NodeSelector = utils.UnionMaps(podSpec.NodeSelector, config.GetK8sPluginConfig().InterruptibleNodeSelector)
8881
}
8982
if podSpec.Affinity == nil && config.GetK8sPluginConfig().DefaultAffinity != nil {
@@ -98,16 +91,11 @@ func UpdatePodWithInterruptibleFlag(taskExecutionMetadata pluginsCore.TaskExecut
9891
if podSpec.DNSConfig == nil && config.GetK8sPluginConfig().DefaultPodDNSConfig != nil {
9992
podSpec.DNSConfig = config.GetK8sPluginConfig().DefaultPodDNSConfig.DeepCopy()
10093
}
101-
ApplyInterruptibleNodeAffinity(isInterruptible, podSpec)
94+
ApplyInterruptibleNodeAffinity(taskExecutionMetadata.IsInterruptible(), podSpec)
10295
}
10396

10497
// ToK8sPodSpec constructs a pod spec from the given TaskTemplate
10598
func ToK8sPodSpec(ctx context.Context, tCtx pluginsCore.TaskExecutionContext) (*v1.PodSpec, error) {
106-
return ToK8sPodSpecWithInterruptible(ctx, tCtx, false)
107-
}
108-
109-
// ToK8sPodSpecWithInterruptible constructs a pod spec from the gien TaskTemplate and optionally add (interruptible instance) support.
110-
func ToK8sPodSpecWithInterruptible(ctx context.Context, tCtx pluginsCore.TaskExecutionContext, omitInterruptible bool) (*v1.PodSpec, error) {
11199
task, err := tCtx.TaskReader().Read(ctx)
112100
if err != nil {
113101
logger.Warnf(ctx, "failed to read task information when trying to construct Pod, err: %s", err.Error())
@@ -138,7 +126,7 @@ func ToK8sPodSpecWithInterruptible(ctx context.Context, tCtx pluginsCore.TaskExe
138126
pod := &v1.PodSpec{
139127
Containers: containers,
140128
}
141-
UpdatePodWithInterruptibleFlag(tCtx.TaskExecutionMetadata(), []v1.ResourceRequirements{c.Resources}, pod, omitInterruptible)
129+
UpdatePod(tCtx.TaskExecutionMetadata(), []v1.ResourceRequirements{c.Resources}, pod)
142130

143131
if err := AddCoPilotToPod(ctx, config.GetK8sPluginConfig().CoPilot, pod, task.GetInterface(), tCtx.TaskExecutionMetadata(), tCtx.InputReader(), tCtx.OutputWriter(), task.GetContainer().GetDataConfig()); err != nil {
144132
return nil, err

go/tasks/pluginmachinery/flytek8s/pod_helper_test.go

Lines changed: 0 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -109,7 +109,6 @@ func TestPodSetup(t *testing.T) {
109109
t.Run("ApplyInterruptibleNodeAffinity", TestApplyInterruptibleNodeAffinity)
110110
t.Run("UpdatePod", updatePod)
111111
t.Run("ToK8sPodInterruptible", toK8sPodInterruptible)
112-
t.Run("toK8sPodInterruptibleFalse", toK8sPodInterruptibleFalse)
113112
}
114113

115114
func TestApplyInterruptibleNodeAffinity(t *testing.T) {
@@ -349,43 +348,6 @@ func toK8sPodInterruptible(t *testing.T) {
349348
)
350349
}
351350

352-
func toK8sPodInterruptibleFalse(t *testing.T) {
353-
ctx := context.TODO()
354-
355-
x := dummyExecContext(&v1.ResourceRequirements{
356-
Limits: v1.ResourceList{
357-
v1.ResourceCPU: resource.MustParse("1024m"),
358-
v1.ResourceStorage: resource.MustParse("100M"),
359-
ResourceNvidiaGPU: resource.MustParse("1"),
360-
},
361-
Requests: v1.ResourceList{
362-
v1.ResourceCPU: resource.MustParse("1024m"),
363-
v1.ResourceStorage: resource.MustParse("100M"),
364-
},
365-
})
366-
367-
p, err := ToK8sPodSpecWithInterruptible(ctx, x, true)
368-
assert.NoError(t, err)
369-
assert.Len(t, p.Tolerations, 1)
370-
assert.Equal(t, 0, len(p.NodeSelector))
371-
assert.Equal(t, "", p.NodeSelector["x/interruptible"])
372-
assert.NotEqualValues(
373-
t,
374-
[]v1.NodeSelectorTerm{
375-
v1.NodeSelectorTerm{
376-
MatchExpressions: []v1.NodeSelectorRequirement{
377-
v1.NodeSelectorRequirement{
378-
Key: "x/interruptible",
379-
Operator: v1.NodeSelectorOpIn,
380-
Values: []string{"true"},
381-
},
382-
},
383-
},
384-
},
385-
p.Affinity.NodeAffinity.RequiredDuringSchedulingIgnoredDuringExecution.NodeSelectorTerms,
386-
)
387-
}
388-
389351
func TestToK8sPod(t *testing.T) {
390352
ctx := context.TODO()
391353

go/tasks/plugins/array/core/state.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,9 @@ type State struct {
5353

5454
// Tracks the number of subtask retries using the execution index
5555
RetryAttempts bitarray.CompactArray `json:"retryAttempts"`
56+
57+
// Tracks the number of system failures for each subtask using the execution index
58+
SystemFailures bitarray.CompactArray `json:"systemFailures"`
5659
}
5760

5861
func (s State) GetReason() string {

go/tasks/plugins/array/k8s/management.go

Lines changed: 39 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -93,7 +93,7 @@ func LaunchAndCheckSubTasksState(ctx context.Context, tCtx core.TaskExecutionCon
9393

9494
retryAttemptsArray, err := bitarray.NewCompactArray(count, maxValue)
9595
if err != nil {
96-
logger.Errorf(context.Background(), "Failed to create attempts compact array with [count: %v, maxValue: %v]", count, maxValue)
96+
logger.Errorf(ctx, "Failed to create attempts compact array with [count: %v, maxValue: %v]", count, maxValue)
9797
return currentState, externalResources, nil
9898
}
9999

@@ -106,6 +106,26 @@ func LaunchAndCheckSubTasksState(ctx context.Context, tCtx core.TaskExecutionCon
106106
currentState.RetryAttempts = retryAttemptsArray
107107
}
108108

109+
// If the current State is newly minted then we must initialize SystemFailures to track how many
110+
// times the subtask failed due to system issues, this is necessary to correctly evaluate
111+
// interruptible subtasks.
112+
if len(currentState.SystemFailures.GetItems()) == 0 {
113+
count := uint(currentState.GetExecutionArraySize())
114+
maxValue := bitarray.Item(tCtx.TaskExecutionMetadata().GetInterruptibleFailureThreshold())
115+
116+
systemFailuresArray, err := bitarray.NewCompactArray(count, maxValue)
117+
if err != nil {
118+
logger.Errorf(ctx, "Failed to create system failures array with [count: %v, maxValue: %v]", count, maxValue)
119+
return currentState, externalResources, err
120+
}
121+
122+
for i := 0; i < currentState.GetExecutionArraySize(); i++ {
123+
systemFailuresArray.SetItem(i, 0)
124+
}
125+
126+
currentState.SystemFailures = systemFailuresArray
127+
}
128+
109129
// initialize log plugin
110130
logPlugin, err := logs.InitializeLogPlugins(&config.LogConfig.Config)
111131
if err != nil {
@@ -146,7 +166,8 @@ func LaunchAndCheckSubTasksState(ctx context.Context, tCtx core.TaskExecutionCon
146166
}
147167

148168
originalIdx := arrayCore.CalculateOriginalIndex(childIdx, newState.GetIndexesToCache())
149-
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, childIdx, originalIdx, retryAttempt)
169+
systemFailures := currentState.SystemFailures.GetItem(childIdx)
170+
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, childIdx, originalIdx, retryAttempt, systemFailures)
150171
if err != nil {
151172
return currentState, externalResources, err
152173
}
@@ -188,6 +209,16 @@ func LaunchAndCheckSubTasksState(ctx context.Context, tCtx core.TaskExecutionCon
188209
return currentState, externalResources, perr
189210
}
190211

212+
if phaseInfo.Err() != nil {
213+
messageCollector.Collect(childIdx, phaseInfo.Err().String())
214+
}
215+
216+
if phaseInfo.Err() != nil && phaseInfo.Err().GetKind() == idlCore.ExecutionError_SYSTEM {
217+
newState.SystemFailures.SetItem(childIdx, systemFailures+1)
218+
} else {
219+
newState.SystemFailures.SetItem(childIdx, systemFailures)
220+
}
221+
191222
// process subtask phase
192223
actualPhase := phaseInfo.Phase()
193224
if actualPhase.IsSuccess() {
@@ -294,15 +325,19 @@ func TerminateSubTasks(ctx context.Context, tCtx core.TaskExecutionContext, kube
294325
messageCollector := errorcollector.NewErrorMessageCollector()
295326
for childIdx, existingPhaseIdx := range currentState.GetArrayStatus().Detailed.GetItems() {
296327
existingPhase := core.Phases[existingPhaseIdx]
297-
retryAttempt := currentState.RetryAttempts.GetItem(childIdx)
328+
retryAttempt := uint64(0)
329+
if childIdx < len(currentState.RetryAttempts.GetItems()) {
330+
// we can use RetryAttempts if it has been initialized, otherwise stay with default 0
331+
retryAttempt = currentState.RetryAttempts.GetItem(childIdx)
332+
}
298333

299334
// return immediately if subtask has completed or not yet started
300335
if existingPhase.IsTerminal() || existingPhase == core.PhaseUndefined {
301336
continue
302337
}
303338

304339
originalIdx := arrayCore.CalculateOriginalIndex(childIdx, currentState.GetIndexesToCache())
305-
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, childIdx, originalIdx, retryAttempt)
340+
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, childIdx, originalIdx, retryAttempt, 0)
306341
if err != nil {
307342
return err
308343
}

go/tasks/plugins/array/k8s/management_test.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -91,6 +91,7 @@ func getMockTaskExecutionContext(ctx context.Context, parallelism int) *mocks.Ta
9191
tMeta.OnGetAnnotations().Return(nil)
9292
tMeta.OnGetOwnerReference().Return(metav1.OwnerReference{})
9393
tMeta.OnGetPlatformResources().Return(&v1.ResourceRequirements{})
94+
tMeta.OnGetInterruptibleFailureThreshold().Return(2)
9495

9596
ow := &mocks2.OutputWriter{}
9697
ow.OnGetOutputPrefixPath().Return("/prefix/")

go/tasks/plugins/array/k8s/subtask_exec_context.go

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -43,9 +43,9 @@ func (s SubTaskExecutionContext) TaskReader() pluginsCore.TaskReader {
4343

4444
// NewSubtaskExecutionContext constructs a SubTaskExecutionContext using the provided parameters
4545
func NewSubTaskExecutionContext(tCtx pluginsCore.TaskExecutionContext, taskTemplate *core.TaskTemplate,
46-
executionIndex, originalIndex int, retryAttempt uint64) (SubTaskExecutionContext, error) {
46+
executionIndex, originalIndex int, retryAttempt uint64, systemFailures uint64) (SubTaskExecutionContext, error) {
4747

48-
subTaskExecutionMetadata, err := NewSubTaskExecutionMetadata(tCtx.TaskExecutionMetadata(), taskTemplate, executionIndex, retryAttempt)
48+
subTaskExecutionMetadata, err := NewSubTaskExecutionMetadata(tCtx.TaskExecutionMetadata(), taskTemplate, executionIndex, retryAttempt, systemFailures)
4949
if err != nil {
5050
return SubTaskExecutionContext{}, err
5151
}
@@ -135,6 +135,7 @@ type SubTaskExecutionMetadata struct {
135135
pluginsCore.TaskExecutionMetadata
136136
annotations map[string]string
137137
labels map[string]string
138+
interruptible bool
138139
subtaskExecutionID SubTaskExecutionID
139140
}
140141

@@ -153,8 +154,14 @@ func (s SubTaskExecutionMetadata) GetTaskExecutionID() pluginsCore.TaskExecution
153154
return s.subtaskExecutionID
154155
}
155156

157+
// IsInterruptbile overrides the base NodeExecutionMetadata to return a subtask specific identifier
158+
func (s SubTaskExecutionMetadata) IsInterruptible() bool {
159+
return s.interruptible
160+
}
161+
156162
// NewSubtaskExecutionMetadata constructs a SubTaskExecutionMetadata using the provided parameters
157-
func NewSubTaskExecutionMetadata(taskExecutionMetadata pluginsCore.TaskExecutionMetadata, taskTemplate *core.TaskTemplate, executionIndex int, retryAttempt uint64) (SubTaskExecutionMetadata, error) {
163+
func NewSubTaskExecutionMetadata(taskExecutionMetadata pluginsCore.TaskExecutionMetadata, taskTemplate *core.TaskTemplate,
164+
executionIndex int, retryAttempt uint64, systemFailures uint64) (SubTaskExecutionMetadata, error) {
158165

159166
var err error
160167
secretsMap := make(map[string]string)
@@ -171,10 +178,12 @@ func NewSubTaskExecutionMetadata(taskExecutionMetadata pluginsCore.TaskExecution
171178
}
172179

173180
subTaskExecutionID := NewSubTaskExecutionID(taskExecutionMetadata.GetTaskExecutionID(), executionIndex, retryAttempt)
181+
interruptible := taskExecutionMetadata.IsInterruptible() && uint32(systemFailures) < taskExecutionMetadata.GetInterruptibleFailureThreshold()
174182
return SubTaskExecutionMetadata{
175183
taskExecutionMetadata,
176184
utils.UnionMaps(taskExecutionMetadata.GetAnnotations(), secretsMap),
177185
utils.UnionMaps(taskExecutionMetadata.GetLabels(), injectSecretsLabel),
186+
interruptible,
178187
subTaskExecutionID,
179188
}, nil
180189
}

go/tasks/plugins/array/k8s/subtask_exec_context_test.go

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,9 @@ func TestSubTaskExecutionContext(t *testing.T) {
2020
executionIndex := 0
2121
originalIndex := 5
2222
retryAttempt := uint64(1)
23+
systemFailures := uint64(0)
2324

24-
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, executionIndex, originalIndex, retryAttempt)
25+
stCtx, err := NewSubTaskExecutionContext(tCtx, taskTemplate, executionIndex, originalIndex, retryAttempt, systemFailures)
2526
assert.Nil(t, err)
2627

2728
assert.Equal(t, fmt.Sprintf("notfound-%d-%d", executionIndex, retryAttempt), stCtx.TaskExecutionMetadata().GetTaskExecutionID().GetGeneratedName())

tests/end_to_end.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -176,6 +176,7 @@ func RunPluginEndToEndTest(t *testing.T, executor pluginCore.Plugin, template *i
176176
Name: execID,
177177
})
178178
tMeta.OnGetPlatformResources().Return(&v1.ResourceRequirements{})
179+
tMeta.OnGetInterruptibleFailureThreshold().Return(2)
179180

180181
catClient := &catalogMocks.Client{}
181182
catData := sync.Map{}

0 commit comments

Comments
 (0)