Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions flyteplugins/go/tasks/plugins/k8s/ray/ray.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,11 @@ import (
pluginsCore "github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/core"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/flytek8s"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/flytek8s/config"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/ioutils"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/k8s"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/tasklog"
pluginsUtils "github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/utils"
"github.com/flyteorg/flyte/flytestdlib/logger"
"github.com/flyteorg/flyte/flytestdlib/utils"
)

Expand Down Expand Up @@ -712,6 +714,24 @@ func (plugin rayJobResourceHandler) GetTaskPhase(ctx context.Context, pluginCont
case rayv1.JobDeploymentStatusFailed:
failInfo := fmt.Sprintf("Failed to run Ray job %s with error: [%s] %s", rayJob.Name, rayJob.Status.Reason, rayJob.Status.Message)
phaseInfo, err = pluginsCore.PhaseInfoSystemRetryableFailureWithCleanup(flyteerr.TaskFailedWithError, failInfo, info), nil
if writer := pluginContext.OutputWriter(); writer != nil {
reader := ioutils.NewRemoteFileOutputReader(ctx, pluginContext.DataStore(), writer, 0)
hasError, readErr := reader.IsError(ctx)
if readErr != nil {
logger.Warnf(ctx, "Failed to check Ray task error file; retaining system retry: %v", readErr)
} else if hasError {
taskError, readErr := reader.ReadError(ctx)
if readErr != nil {
logger.Warnf(ctx, "Failed to read Ray task error file; retaining system retry: %v", readErr)
} else if taskError.Kind == core.ExecutionError_USER {
if taskError.IsRecoverable {
phaseInfo = pluginsCore.PhaseInfoRetryableFailureWithCleanup(flyteerr.TaskFailedWithError, failInfo, info)
} else {
phaseInfo = pluginsCore.PhaseInfoFailureWithCleanup(flyteerr.TaskFailedWithError, failInfo, info)
}
}
}
}
default:
// We already handle all known deployment status, so this should never happen unless a future version of ray
// introduced a new job status.
Expand Down
107 changes: 105 additions & 2 deletions flyteplugins/go/tasks/plugins/k8s/ray/ray_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"reflect"
"testing"
"time"
Expand All @@ -26,10 +28,16 @@ import (
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/flytek8s"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/flytek8s/config"
pluginIOMocks "github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/io/mocks"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/ioutils"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/k8s"
mocks2 "github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/k8s/mocks"
"github.com/flyteorg/flyte/flyteplugins/go/tasks/pluginmachinery/tasklog"
"github.com/flyteorg/flyte/flytestdlib/contextutils"
"github.com/flyteorg/flyte/flytestdlib/promutils"
"github.com/flyteorg/flyte/flytestdlib/promutils/labeled"
"github.com/flyteorg/flyte/flytestdlib/storage"
"github.com/flyteorg/flyte/flytestdlib/utils"
"github.com/flyteorg/stow/local"
)

const (
Expand Down Expand Up @@ -1172,7 +1180,8 @@ func TestGetTaskPhaseFailedRetryable(t *testing.T) {
// reason and message from the RayJob status.
ctx := context.Background()
rayJobResourceHandler := rayJobResourceHandler{}
pluginCtx := newPluginContext(k8s.PluginState{})
pluginCtx := newPluginContext(k8s.PluginState{}).(*mocks2.PluginContext)
pluginCtx.EXPECT().OutputWriter().Return(nil).Maybe()

rayObject := &rayv1.RayJob{
ObjectMeta: metav1.ObjectMeta{
Expand All @@ -1191,6 +1200,99 @@ func TestGetTaskPhaseFailedRetryable(t *testing.T) {
assert.Contains(t, phaseInfo.Err().GetMessage(), "head node ran out of memory")
}

func TestGetTaskPhaseTaskError(t *testing.T) {
labeled.SetMetricKeys(contextutils.ExecIDKey)
cases := []struct {
name string
writeError bool
corrupt bool
oversized bool
origin core.ExecutionError_ErrorKind
recoverable bool
expectedPhase pluginsCore.Phase
expectedKind core.ExecutionError_ErrorKind
}{
{
name: "permanent user error", writeError: true, origin: core.ExecutionError_USER,
expectedPhase: pluginsCore.PhasePermanentFailure, expectedKind: core.ExecutionError_USER,
},
{
name: "recoverable user error", writeError: true, origin: core.ExecutionError_USER, recoverable: true,
expectedPhase: pluginsCore.PhaseRetryableFailure, expectedKind: core.ExecutionError_USER,
},
{
name: "system error keeps system retry", writeError: true, origin: core.ExecutionError_SYSTEM,
expectedPhase: pluginsCore.PhaseRetryableFailure, expectedKind: core.ExecutionError_SYSTEM,
},
{
name: "missing error keeps system retry",
expectedPhase: pluginsCore.PhaseRetryableFailure, expectedKind: core.ExecutionError_SYSTEM,
},
{
name: "corrupt error keeps system retry", corrupt: true,
expectedPhase: pluginsCore.PhaseRetryableFailure, expectedKind: core.ExecutionError_SYSTEM,
},
{
name: "oversized error keeps system retry", oversized: true,
expectedPhase: pluginsCore.PhaseRetryableFailure, expectedKind: core.ExecutionError_SYSTEM,
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
directory := t.TempDir()
store, err := storage.NewDataStore(&storage.Config{
Type: storage.TypeLocal,
InitContainer: directory,
MultiContainerEnabled: true,
Stow: storage.StowConfig{
Kind: local.Kind,
Config: map[string]string{local.ConfigKeyPath: "/"},
},
Limits: storage.LimitsConfig{GetLimitMegabytes: 2},
}, promutils.NewTestScope())
require.NoError(t, err)
paths := ioutils.NewReadOnlyOutputFilePaths(ctx, store, storage.DataReference("file://"+directory))
writer := ioutils.NewRemoteFileOutputWriter(ctx, store, paths)
if tc.writeError {
kind := core.ContainerError_NON_RECOVERABLE
if tc.recoverable {
kind = core.ContainerError_RECOVERABLE
}
require.NoError(t, store.WriteProtobuf(ctx, paths.GetErrorPath(), storage.Options{}, &core.ErrorDocument{
Error: &core.ContainerError{Code: "task-error", Message: "task failed", Kind: kind, Origin: tc.origin},
}))
}
if tc.corrupt {
require.NoError(t, os.WriteFile(filepath.Join(directory, ioutils.ErrorsSuffix), []byte{0xff}, 0o600))
}
if tc.oversized {
file, err := os.Create(filepath.Join(directory, ioutils.ErrorsSuffix))
require.NoError(t, err)
require.NoError(t, file.Truncate(storage.GetConfig().Limits.GetLimitMegabytes*storage.MiB+1))
require.NoError(t, file.Close())
}
pluginContext := newPluginContext(k8s.PluginState{}).(*mocks2.PluginContext)
pluginContext.EXPECT().OutputWriter().Return(writer)
pluginContext.EXPECT().DataStore().Return(store)
job := &rayv1.RayJob{
ObjectMeta: metav1.ObjectMeta{Name: "failed-ray-job"},
Status: rayv1.RayJobStatus{
JobDeploymentStatus: rayv1.JobDeploymentStatusFailed,
Reason: rayv1.AppFailed,
Message: "driver exited",
},
}
phase, err := (rayJobResourceHandler{}).GetTaskPhase(ctx, pluginContext, job)
require.NoError(t, err)
assert.Equal(t, tc.expectedPhase, phase.Phase())
require.NotNil(t, phase.Err())
assert.Equal(t, tc.expectedKind, phase.Err().GetKind())
assert.True(t, phase.CleanupOnFailure())
})
}
}

func newPluginContext(pluginState k8s.PluginState) k8s.PluginContext {
plg := &mocks2.PluginContext{}

Expand Down Expand Up @@ -1252,7 +1354,8 @@ func init() {
func TestGetTaskPhase(t *testing.T) {
ctx := context.Background()
rayJobResourceHandler := rayJobResourceHandler{}
pluginCtx := newPluginContext(k8s.PluginState{})
pluginCtx := newPluginContext(k8s.PluginState{}).(*mocks2.PluginContext)
pluginCtx.EXPECT().OutputWriter().Return(nil).Maybe()

testCases := []struct {
rayJobPhase rayv1.JobDeploymentStatus
Expand Down
Loading