From fb7347d9164483d3d179862ab9efbc46cc7c1de6 Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Thu, 3 Sep 2026 22:53:49 +0530 Subject: [PATCH 1/7] FEAT: Add PySpark bundle and script_uri support to sparkeks. Enable CI-published zip bundles and ingest-style single-file submits via parameters.script_uri, with command-level image override and safer Ranger extension merging. --- .../pkg/object/command/sparkeks/entrypoint.go | 86 ++++++++++ .../command/sparkeks/entrypoint_test.go | 151 ++++++++++++++++++ .../pkg/object/command/sparkeks/sparkeks.go | 107 ++++++++----- .../object/command/sparkeks/sparkeks_test.go | 20 +++ 4 files changed, 325 insertions(+), 39 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index f5aebb70..be00b74f 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -76,6 +76,63 @@ func (s sqlWrapperEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) return nil } +type pysparkEntrypointStrategy struct { + appName string + queryURI string + user string + resultURI string + returnResult bool + arguments []string + scriptURI string // pre-uploaded .py; job passes parameters.script_uri + bundleURI string // command bundle_uri; job never sets this + bundleVersion string // job passes parameters.bundle_version + entryPoint string // job passes parameters.entry_point (path inside the zip) +} + +func (s pysparkEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) error { + if s.scriptURI != "" { + extra := append([]string{s.scriptURI, ""}, s.arguments...) + spec.Arguments = buildArguments(extra, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) + return nil + } + + if strings.TrimSpace(s.bundleVersion) == "" { + return ErrMissingBundleVersion + } + entryPoint := strings.TrimSpace(s.entryPoint) + if entryPoint == "" { + return ErrMissingBundleEntry + } + + bundleZipURI := updateS3ToS3aURI(strings.TrimRight(s.bundleURI, "/") + "/" + strings.TrimSpace(s.bundleVersion) + ".zip") + extra := append([]string{bundleZipURI, entryPoint}, s.arguments...) + spec.Arguments = buildArguments(extra, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) + return nil +} + +func validateScriptURI(scriptURI, allowedPrefix string) error { + scriptURI = strings.TrimSpace(scriptURI) + if scriptURI == "" { + return ErrInvalidScriptURI + } + if !strings.HasPrefix(scriptURI, s3Prefix) && !strings.HasPrefix(scriptURI, s3aPrefix) { + return ErrInvalidScriptURI + } + if strings.Contains(scriptURI, "..") { + return ErrInvalidScriptURI + } + if !strings.HasSuffix(strings.ToLower(scriptURI), ".py") { + return ErrInvalidScriptURI + } + + normalized := updateS3ToS3aURI(scriptURI) + prefix := updateS3ToS3aURI(strings.TrimRight(strings.TrimSpace(allowedPrefix), "/")) + "/" + if !strings.HasPrefix(normalized, prefix) { + return ErrInvalidScriptURI + } + return nil +} + // entrypointFactory builds the entrypoint strategy for a job from its execution context. type entrypointFactory func(execCtx *executionContext) entrypointStrategy @@ -87,6 +144,10 @@ var entrypointStrategiesByExt = map[string]entrypointFactory{ var defaultEntrypointFactory entrypointFactory = newSQLWrapperEntrypointStrategy func newEntrypointStrategy(execCtx *executionContext) entrypointStrategy { + if execCtx.commandContext.BundleURI != "" { + return newPySparkEntrypointStrategy(execCtx) + } + ext := strings.ToLower(path.Ext(execCtx.commandContext.WrapperURI)) factory, ok := entrypointStrategiesByExt[ext] if !ok { @@ -95,6 +156,31 @@ func newEntrypointStrategy(execCtx *executionContext) entrypointStrategy { return factory(execCtx) } +func newPySparkEntrypointStrategy(execCtx *executionContext) entrypointStrategy { + jobContext := execCtx.jobContext + + s := pysparkEntrypointStrategy{ + appName: execCtx.appName, + queryURI: execCtx.s3aQueryURI, + user: execCtx.job.User, + resultURI: execCtx.s3aResultURI, + returnResult: jobContext.ReturnResult, + arguments: jobContext.Arguments, + bundleURI: execCtx.commandContext.BundleURI, + } + + if jobContext.Parameters != nil && strings.TrimSpace(jobContext.Parameters.ScriptURI) != "" { + s.scriptURI = updateS3ToS3aURI(strings.TrimSpace(jobContext.Parameters.ScriptURI)) + return s + } + + if jobContext.Parameters != nil { + s.entryPoint = jobContext.Parameters.EntryPoint + s.bundleVersion = jobContext.Parameters.BundleVersion + } + return s +} + func newJarEntrypointStrategy(execCtx *executionContext) entrypointStrategy { jobContext := execCtx.jobContext diff --git a/internal/pkg/object/command/sparkeks/entrypoint_test.go b/internal/pkg/object/command/sparkeks/entrypoint_test.go index 17db798a..0bbd9518 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint_test.go +++ b/internal/pkg/object/command/sparkeks/entrypoint_test.go @@ -189,3 +189,154 @@ func newTestExecutionContext(wrapperURI string, jobCtx *jobContext, user string, s3aResultURI: resultURI, } } + +func newTestBundleExecutionContext(bundleURI string, jobCtx *jobContext, user string, resultURI string) *executionContext { + execCtx := newTestExecutionContext("s3://bucket/pyspark-job-wrapper.py", jobCtx, user, resultURI) + execCtx.commandContext.BundleURI = bundleURI + return execCtx +} + +func TestNewEntrypointStrategy_BundleURISelectsPySparkStrategy(t *testing.T) { + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/buybox-predictor/", &jobContext{}, "alice", "result_uri") + strategy := newEntrypointStrategy(execCtx) + if got := reflect.TypeOf(strategy); got != reflect.TypeOf(pysparkEntrypointStrategy{}) { + t.Errorf("newEntrypointStrategy() type = %v, want pysparkEntrypointStrategy", got) + } +} + +func TestNewEntrypointStrategy_ScriptURISelectsPySparkStrategy(t *testing.T) { + jobCtx := &jobContext{ + Parameters: &jobParameters{ + ScriptURI: "s3://bucket/pyspark/scripts/alice/train.py", + }, + } + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "result_uri") + strategy := newEntrypointStrategy(execCtx) + if got := reflect.TypeOf(strategy); got != reflect.TypeOf(pysparkEntrypointStrategy{}) { + t.Errorf("newEntrypointStrategy() type = %v, want pysparkEntrypointStrategy", got) + } +} + +func TestPySparkBundleEntrypointStrategy_Apply(t *testing.T) { + jobCtx := &jobContext{ + ReturnResult: true, + Arguments: []string{"--date=2026-08-31"}, + Parameters: &jobParameters{ + EntryPoint: "src/prediction_flow.py", + BundleVersion: "2f3a91c", + }, + } + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/buybox-predictor/", jobCtx, "alice", "result_uri") + + strategy := newEntrypointStrategy(execCtx) + spec := &v1beta2.SparkApplicationSpec{} + if err := strategy.apply(spec); err != nil { + t.Fatalf("apply() returned unexpected error: %v", err) + } + + expectedArgs := []string{ + "spark-sql-job-test", "s3a://bucket/query.sql", "alice", "result_uri", + "s3a://bucket/pyspark/buybox-predictor/2f3a91c.zip", "src/prediction_flow.py", + "--date=2026-08-31", + } + if !reflect.DeepEqual(spec.Arguments, expectedArgs) { + t.Errorf("spec.Arguments = %v, want %v", spec.Arguments, expectedArgs) + } +} + +func TestPySparkScriptEntrypointStrategy_Apply(t *testing.T) { + jobCtx := &jobContext{ + ReturnResult: true, + Arguments: []string{"--date=2026-08-31"}, + Parameters: &jobParameters{ + ScriptURI: "s3://bucket/pyspark/scripts/alice/train.py", + }, + } + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "result_uri") + + strategy := newEntrypointStrategy(execCtx) + spec := &v1beta2.SparkApplicationSpec{} + if err := strategy.apply(spec); err != nil { + t.Fatalf("apply() returned unexpected error: %v", err) + } + + expectedArgs := []string{ + "spark-sql-job-test", "s3a://bucket/query.sql", "alice", "result_uri", + "s3a://bucket/pyspark/scripts/alice/train.py", "", + "--date=2026-08-31", + } + if !reflect.DeepEqual(spec.Arguments, expectedArgs) { + t.Errorf("spec.Arguments = %v, want %v", spec.Arguments, expectedArgs) + } +} + +func TestValidateScriptURI(t *testing.T) { + prefix := "s3://bucket/pyspark/" + tests := []struct { + name string + scriptURI string + wantErr error + }{ + {name: "valid", scriptURI: "s3://bucket/pyspark/scripts/alice/train.py"}, + {name: "valid s3a", scriptURI: "s3a://bucket/pyspark/scripts/alice/train.py"}, + {name: "wrong prefix", scriptURI: "s3://other-bucket/pyspark/train.py", wantErr: ErrInvalidScriptURI}, + {name: "not python", scriptURI: "s3://bucket/pyspark/scripts/alice/train.sh", wantErr: ErrInvalidScriptURI}, + {name: "traversal", scriptURI: "s3://bucket/pyspark/../evil.py", wantErr: ErrInvalidScriptURI}, + {name: "not s3", scriptURI: "https://bucket/pyspark/train.py", wantErr: ErrInvalidScriptURI}, + {name: "empty", scriptURI: "", wantErr: ErrInvalidScriptURI}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateScriptURI(tt.scriptURI, prefix) + if tt.wantErr == nil && err != nil { + t.Fatalf("validateScriptURI() error = %v, want nil", err) + } + if tt.wantErr != nil && err != tt.wantErr { + t.Errorf("validateScriptURI() error = %v, want %v", err, tt.wantErr) + } + }) + } +} + +func TestPySparkBundleEntrypointStrategy_Apply_MissingFields(t *testing.T) { + tests := []struct { + name string + jobCtx *jobContext + wantErr error + }{ + {name: "nil parameters", jobCtx: &jobContext{}, wantErr: ErrMissingBundleVersion}, + {name: "missing bundle_version", jobCtx: &jobContext{Parameters: &jobParameters{EntryPoint: "src/main.py"}}, wantErr: ErrMissingBundleVersion}, + {name: "missing entry_point", jobCtx: &jobContext{Parameters: &jobParameters{BundleVersion: "abc123"}}, wantErr: ErrMissingBundleEntry}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/proj/", tt.jobCtx, "alice", "result_uri") + strategy := newEntrypointStrategy(execCtx) + spec := &v1beta2.SparkApplicationSpec{} + if err := strategy.apply(spec); err != tt.wantErr { + t.Errorf("apply() error = %v, want %v", err, tt.wantErr) + } + }) + } +} + +func TestPySparkBundleEntrypointStrategy_Apply_ProjectSlashVersion(t *testing.T) { + jobCtx := &jobContext{ + Parameters: &jobParameters{ + EntryPoint: "src/update_prediction_data.py", + BundleVersion: "buybox-predictor/2f3a91c", + }, + } + execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "") + + strategy := newEntrypointStrategy(execCtx) + spec := &v1beta2.SparkApplicationSpec{} + if err := strategy.apply(spec); err != nil { + t.Fatalf("apply() returned unexpected error: %v", err) + } + if spec.Arguments[4] != "s3a://bucket/pyspark/buybox-predictor/2f3a91c.zip" { + t.Errorf("bundle uri = %q", spec.Arguments[4]) + } +} diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index 6df2f2e0..b222766b 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -95,11 +95,17 @@ var ( ErrApplicationSpec = fmt.Errorf("failed to load or parse SparkApplication template") ErrSparkApplicationFile = fmt.Errorf("failed to read SparkApplication application template file: check file path and permissions") ErrMissingEntryPoint = fmt.Errorf("entry_point is required for .jar wrapper_uri: set parameters.entry_point to the fully-qualified main class") + ErrMissingBundleVersion = fmt.Errorf("parameters.bundle_version is required when the command sets bundle_uri") + ErrMissingBundleEntry = fmt.Errorf("parameters.entry_point is required when the command sets bundle_uri") + ErrInvalidScriptURI = fmt.Errorf("parameters.script_uri must be an s3:// or s3a:// .py object under the command bundle_uri prefix") + ErrConflictingPySparkSource = fmt.Errorf("parameters.script_uri and parameters.bundle_version are mutually exclusive") ) type commandContext struct { - JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` - WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_uri,omitempty"` + JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` + WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_uri,omitempty"` + BundleURI string `yaml:"bundle_uri,omitempty" json:"bundle_uri,omitempty"` + Image string `yaml:"image,omitempty" json:"image,omitempty"` EventLogURI string `yaml:"event_log_uri,omitempty" json:"event_log_uri,omitempty"` Properties map[string]string `yaml:"properties,omitempty" json:"properties,omitempty"` KubeNamespace string `yaml:"kube_namespace,omitempty" json:"kube_namespace,omitempty"` @@ -109,6 +115,8 @@ type jobParameters struct { Properties map[string]string `yaml:"properties,omitempty" json:"properties,omitempty"` EntryPoint string `yaml:"entry_point,omitempty" json:"entry_point,omitempty"` ApplicationType string `yaml:"application_type,omitempty" json:"application_type,omitempty"` + BundleVersion string `yaml:"bundle_version,omitempty" json:"bundle_version,omitempty"` + ScriptURI string `yaml:"script_uri,omitempty" json:"script_uri,omitempty"` } type jobContext struct { @@ -340,6 +348,19 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. } execCtx.jobContext = jobContext + if s.BundleURI != "" && jobContext.Parameters != nil { + scriptURI := strings.TrimSpace(jobContext.Parameters.ScriptURI) + bundleVersion := strings.TrimSpace(jobContext.Parameters.BundleVersion) + if scriptURI != "" && bundleVersion != "" { + return nil, ErrConflictingPySparkSource + } + if scriptURI != "" { + if err := validateScriptURI(scriptURI, s.BundleURI); err != nil { + return nil, err + } + } + } + // Parse cluster context clusterContext := &clusterContext{} if c.Context != nil { @@ -727,6 +748,17 @@ func updateKubeConfig(ctx context.Context, execCtx *executionContext) (string, e return kubeconfigPath, nil } +// imageForJob prefers the command-pinned image (DS CI) over the cluster default (SQL image). +func imageForJob(cmd *commandContext, cluster *clusterContext) *string { + if cmd != nil && cmd.Image != "" { + return &cmd.Image + } + if cluster != nil { + return cluster.Image + } + return nil +} + // applySparkOperatorConfig consolidates all Spark Operator configuration updates and overrides. func applySparkOperatorConfig(execCtx *executionContext) error { sparkApp := execCtx.sparkApp @@ -766,8 +798,8 @@ func applySparkOperatorConfig(execCtx *executionContext) error { } } - if clusterContext.Image != nil { - sparkApp.Spec.Image = clusterContext.Image + if img := imageForJob(execCtx.commandContext, clusterContext); img != nil { + sparkApp.Spec.Image = img } if clusterContext.Region != nil { @@ -777,41 +809,6 @@ func applySparkOperatorConfig(execCtx *executionContext) error { sparkApp.Spec.Driver.EnvVars[awsRegionEnvVar] = *clusterContext.Region } - // Handle required Spark SQL extensions - if clusterContext.RequiredSparkSQLExtensions != "" { - existingExtensions := sparkApp.Spec.SparkConf[sparkSqlExtensions] - if existingExtensions == "" { - // No existing extensions, just set the required ones - sparkApp.Spec.SparkConf[sparkSqlExtensions] = clusterContext.RequiredSparkSQLExtensions - } else { - // Merge required extensions with existing ones, avoiding duplicates - extensionSet := make(map[string]bool) - - // Add existing extensions to the set - for _, ext := range strings.Split(existingExtensions, ",") { - ext = strings.TrimSpace(ext) - if ext != "" { - extensionSet[ext] = true - } - } - - // Add required extensions to the set - for _, ext := range strings.Split(clusterContext.RequiredSparkSQLExtensions, ",") { - ext = strings.TrimSpace(ext) - if ext != "" { - extensionSet[ext] = true - } - } - - // Build the final extension list - var extensions []string - for ext := range extensionSet { - extensions = append(extensions, ext) - } - sparkApp.Spec.SparkConf[sparkSqlExtensions] = strings.Join(extensions, ",") - } - } - // Driver and Executor resources are handled by deleting from job properties after use // to avoid them being added to sparkConf directly. if driverCores := jobContext.Parameters.Properties[sparkDriverCoresKey]; driverCores != "" { @@ -851,6 +848,38 @@ func applySparkOperatorConfig(execCtx *executionContext) error { sparkApp.Spec.SparkConf[k] = v } + // Required Spark SQL extensions must win over job/cluster properties, so this + // merge runs last: a caller can otherwise submit spark.sql.extensions="" and + // disable the Ranger authorization extension entirely. + if clusterContext.RequiredSparkSQLExtensions != "" { + existingExtensions := sparkApp.Spec.SparkConf[sparkSqlExtensions] + if existingExtensions == "" { + sparkApp.Spec.SparkConf[sparkSqlExtensions] = clusterContext.RequiredSparkSQLExtensions + } else { + extensionSet := make(map[string]bool) + + for _, ext := range strings.Split(existingExtensions, ",") { + ext = strings.TrimSpace(ext) + if ext != "" { + extensionSet[ext] = true + } + } + + for _, ext := range strings.Split(clusterContext.RequiredSparkSQLExtensions, ",") { + ext = strings.TrimSpace(ext) + if ext != "" { + extensionSet[ext] = true + } + } + + var extensions []string + for ext := range extensionSet { + extensions = append(extensions, ext) + } + sparkApp.Spec.SparkConf[sparkSqlExtensions] = strings.Join(extensions, ",") + } + } + if sparkApp.Spec.Type == "" { sparkApp.Spec.Type = defaultApplicationType } diff --git a/internal/pkg/object/command/sparkeks/sparkeks_test.go b/internal/pkg/object/command/sparkeks/sparkeks_test.go index dc823c4c..015b38c9 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks_test.go +++ b/internal/pkg/object/command/sparkeks/sparkeks_test.go @@ -87,3 +87,23 @@ func TestPrintState(t *testing.T) { printState(f, "RUNNING") // Optionally check file contents if needed } + +func TestImageForJob(t *testing.T) { + clusterImg := "cluster.example/spark:v4.1.1" + cluster := &clusterContext{Image: &clusterImg} + + got := imageForJob(&commandContext{Image: "ds.example/spark-ds:v4.1.1"}, cluster) + if got == nil || *got != "ds.example/spark-ds:v4.1.1" { + t.Fatalf("command image should win, got %v", got) + } + + got = imageForJob(&commandContext{}, cluster) + if got == nil || *got != clusterImg { + t.Fatalf("empty command image should fall back to cluster, got %v", got) + } + + got = imageForJob(&commandContext{}, &clusterContext{}) + if got != nil { + t.Fatalf("no images set, got %v", got) + } +} From 58d40f2da9b7abafef53f775dadc305d323b4648 Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Thu, 3 Sep 2026 22:54:56 +0530 Subject: [PATCH 2/7] Drop sparkeks test additions from PySpark entrypoint work. --- .../command/sparkeks/entrypoint_test.go | 151 ------------------ .../object/command/sparkeks/sparkeks_test.go | 20 --- 2 files changed, 171 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint_test.go b/internal/pkg/object/command/sparkeks/entrypoint_test.go index 0bbd9518..17db798a 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint_test.go +++ b/internal/pkg/object/command/sparkeks/entrypoint_test.go @@ -189,154 +189,3 @@ func newTestExecutionContext(wrapperURI string, jobCtx *jobContext, user string, s3aResultURI: resultURI, } } - -func newTestBundleExecutionContext(bundleURI string, jobCtx *jobContext, user string, resultURI string) *executionContext { - execCtx := newTestExecutionContext("s3://bucket/pyspark-job-wrapper.py", jobCtx, user, resultURI) - execCtx.commandContext.BundleURI = bundleURI - return execCtx -} - -func TestNewEntrypointStrategy_BundleURISelectsPySparkStrategy(t *testing.T) { - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/buybox-predictor/", &jobContext{}, "alice", "result_uri") - strategy := newEntrypointStrategy(execCtx) - if got := reflect.TypeOf(strategy); got != reflect.TypeOf(pysparkEntrypointStrategy{}) { - t.Errorf("newEntrypointStrategy() type = %v, want pysparkEntrypointStrategy", got) - } -} - -func TestNewEntrypointStrategy_ScriptURISelectsPySparkStrategy(t *testing.T) { - jobCtx := &jobContext{ - Parameters: &jobParameters{ - ScriptURI: "s3://bucket/pyspark/scripts/alice/train.py", - }, - } - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "result_uri") - strategy := newEntrypointStrategy(execCtx) - if got := reflect.TypeOf(strategy); got != reflect.TypeOf(pysparkEntrypointStrategy{}) { - t.Errorf("newEntrypointStrategy() type = %v, want pysparkEntrypointStrategy", got) - } -} - -func TestPySparkBundleEntrypointStrategy_Apply(t *testing.T) { - jobCtx := &jobContext{ - ReturnResult: true, - Arguments: []string{"--date=2026-08-31"}, - Parameters: &jobParameters{ - EntryPoint: "src/prediction_flow.py", - BundleVersion: "2f3a91c", - }, - } - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/buybox-predictor/", jobCtx, "alice", "result_uri") - - strategy := newEntrypointStrategy(execCtx) - spec := &v1beta2.SparkApplicationSpec{} - if err := strategy.apply(spec); err != nil { - t.Fatalf("apply() returned unexpected error: %v", err) - } - - expectedArgs := []string{ - "spark-sql-job-test", "s3a://bucket/query.sql", "alice", "result_uri", - "s3a://bucket/pyspark/buybox-predictor/2f3a91c.zip", "src/prediction_flow.py", - "--date=2026-08-31", - } - if !reflect.DeepEqual(spec.Arguments, expectedArgs) { - t.Errorf("spec.Arguments = %v, want %v", spec.Arguments, expectedArgs) - } -} - -func TestPySparkScriptEntrypointStrategy_Apply(t *testing.T) { - jobCtx := &jobContext{ - ReturnResult: true, - Arguments: []string{"--date=2026-08-31"}, - Parameters: &jobParameters{ - ScriptURI: "s3://bucket/pyspark/scripts/alice/train.py", - }, - } - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "result_uri") - - strategy := newEntrypointStrategy(execCtx) - spec := &v1beta2.SparkApplicationSpec{} - if err := strategy.apply(spec); err != nil { - t.Fatalf("apply() returned unexpected error: %v", err) - } - - expectedArgs := []string{ - "spark-sql-job-test", "s3a://bucket/query.sql", "alice", "result_uri", - "s3a://bucket/pyspark/scripts/alice/train.py", "", - "--date=2026-08-31", - } - if !reflect.DeepEqual(spec.Arguments, expectedArgs) { - t.Errorf("spec.Arguments = %v, want %v", spec.Arguments, expectedArgs) - } -} - -func TestValidateScriptURI(t *testing.T) { - prefix := "s3://bucket/pyspark/" - tests := []struct { - name string - scriptURI string - wantErr error - }{ - {name: "valid", scriptURI: "s3://bucket/pyspark/scripts/alice/train.py"}, - {name: "valid s3a", scriptURI: "s3a://bucket/pyspark/scripts/alice/train.py"}, - {name: "wrong prefix", scriptURI: "s3://other-bucket/pyspark/train.py", wantErr: ErrInvalidScriptURI}, - {name: "not python", scriptURI: "s3://bucket/pyspark/scripts/alice/train.sh", wantErr: ErrInvalidScriptURI}, - {name: "traversal", scriptURI: "s3://bucket/pyspark/../evil.py", wantErr: ErrInvalidScriptURI}, - {name: "not s3", scriptURI: "https://bucket/pyspark/train.py", wantErr: ErrInvalidScriptURI}, - {name: "empty", scriptURI: "", wantErr: ErrInvalidScriptURI}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - err := validateScriptURI(tt.scriptURI, prefix) - if tt.wantErr == nil && err != nil { - t.Fatalf("validateScriptURI() error = %v, want nil", err) - } - if tt.wantErr != nil && err != tt.wantErr { - t.Errorf("validateScriptURI() error = %v, want %v", err, tt.wantErr) - } - }) - } -} - -func TestPySparkBundleEntrypointStrategy_Apply_MissingFields(t *testing.T) { - tests := []struct { - name string - jobCtx *jobContext - wantErr error - }{ - {name: "nil parameters", jobCtx: &jobContext{}, wantErr: ErrMissingBundleVersion}, - {name: "missing bundle_version", jobCtx: &jobContext{Parameters: &jobParameters{EntryPoint: "src/main.py"}}, wantErr: ErrMissingBundleVersion}, - {name: "missing entry_point", jobCtx: &jobContext{Parameters: &jobParameters{BundleVersion: "abc123"}}, wantErr: ErrMissingBundleEntry}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/proj/", tt.jobCtx, "alice", "result_uri") - strategy := newEntrypointStrategy(execCtx) - spec := &v1beta2.SparkApplicationSpec{} - if err := strategy.apply(spec); err != tt.wantErr { - t.Errorf("apply() error = %v, want %v", err, tt.wantErr) - } - }) - } -} - -func TestPySparkBundleEntrypointStrategy_Apply_ProjectSlashVersion(t *testing.T) { - jobCtx := &jobContext{ - Parameters: &jobParameters{ - EntryPoint: "src/update_prediction_data.py", - BundleVersion: "buybox-predictor/2f3a91c", - }, - } - execCtx := newTestBundleExecutionContext("s3://bucket/pyspark/", jobCtx, "alice", "") - - strategy := newEntrypointStrategy(execCtx) - spec := &v1beta2.SparkApplicationSpec{} - if err := strategy.apply(spec); err != nil { - t.Fatalf("apply() returned unexpected error: %v", err) - } - if spec.Arguments[4] != "s3a://bucket/pyspark/buybox-predictor/2f3a91c.zip" { - t.Errorf("bundle uri = %q", spec.Arguments[4]) - } -} diff --git a/internal/pkg/object/command/sparkeks/sparkeks_test.go b/internal/pkg/object/command/sparkeks/sparkeks_test.go index 015b38c9..dc823c4c 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks_test.go +++ b/internal/pkg/object/command/sparkeks/sparkeks_test.go @@ -87,23 +87,3 @@ func TestPrintState(t *testing.T) { printState(f, "RUNNING") // Optionally check file contents if needed } - -func TestImageForJob(t *testing.T) { - clusterImg := "cluster.example/spark:v4.1.1" - cluster := &clusterContext{Image: &clusterImg} - - got := imageForJob(&commandContext{Image: "ds.example/spark-ds:v4.1.1"}, cluster) - if got == nil || *got != "ds.example/spark-ds:v4.1.1" { - t.Fatalf("command image should win, got %v", got) - } - - got = imageForJob(&commandContext{}, cluster) - if got == nil || *got != clusterImg { - t.Fatalf("empty command image should fall back to cluster, got %v", got) - } - - got = imageForJob(&commandContext{}, &clusterContext{}) - if got != nil { - t.Fatalf("no images set, got %v", got) - } -} From 8c69884150036b3d1ea25a2e66ada0795b94c903 Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Fri, 4 Sep 2026 12:33:15 +0530 Subject: [PATCH 3/7] docs: drop internal image-override wording from sparkeks The command image pin is a generic Spark submit option, not a DS-specific path. --- internal/pkg/object/command/sparkeks/sparkeks.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index b222766b..2ec0bf14 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -748,7 +748,7 @@ func updateKubeConfig(ctx context.Context, execCtx *executionContext) (string, e return kubeconfigPath, nil } -// imageForJob prefers the command-pinned image (DS CI) over the cluster default (SQL image). +// imageForJob prefers a command-pinned image over the cluster default. func imageForJob(cmd *commandContext, cluster *clusterContext) *string { if cmd != nil && cmd.Image != "" { return &cmd.Image From 0e6bc5bc5aea2cf309e820b16cba2773b0ba0f1a Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Tue, 8 Sep 2026 12:20:03 +0530 Subject: [PATCH 4/7] Review + other changes --- .../pkg/object/command/sparkeks/entrypoint.go | 31 ++++++++------ .../pkg/object/command/sparkeks/sparkeks.go | 42 ++++++++++--------- 2 files changed, 41 insertions(+), 32 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index be00b74f..3fea9048 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -83,16 +83,13 @@ type pysparkEntrypointStrategy struct { resultURI string returnResult bool arguments []string - scriptURI string // pre-uploaded .py; job passes parameters.script_uri - bundleURI string // command bundle_uri; job never sets this bundleVersion string // job passes parameters.bundle_version entryPoint string // job passes parameters.entry_point (path inside the zip) } func (s pysparkEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) error { - if s.scriptURI != "" { - extra := append([]string{s.scriptURI, ""}, s.arguments...) - spec.Arguments = buildArguments(extra, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) + if strings.HasSuffix(strings.ToLower(s.queryURI), ".py") { + spec.Arguments = buildArguments(s.arguments, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) return nil } @@ -104,12 +101,24 @@ func (s pysparkEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) err return ErrMissingBundleEntry } - bundleZipURI := updateS3ToS3aURI(strings.TrimRight(s.bundleURI, "/") + "/" + strings.TrimSpace(s.bundleVersion) + ".zip") - extra := append([]string{bundleZipURI, entryPoint}, s.arguments...) + extra := append([]string{entryPoint}, s.arguments...) spec.Arguments = buildArguments(extra, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) return nil } +func pysparkQueryURI(cmd *commandContext, jobCtx *jobContext) string { + if cmd == nil || strings.TrimSpace(cmd.BundleURI) == "" || jobCtx == nil || jobCtx.Parameters == nil { + return "" + } + if script := strings.TrimSpace(jobCtx.Parameters.ScriptURI); script != "" { + return updateS3ToS3aURI(script) + } + if ver := strings.TrimSpace(jobCtx.Parameters.BundleVersion); ver != "" { + return updateS3ToS3aURI(strings.TrimRight(cmd.BundleURI, "/") + "/" + ver + ".zip") + } + return "" +} + func validateScriptURI(scriptURI, allowedPrefix string) error { scriptURI = strings.TrimSpace(scriptURI) if scriptURI == "" { @@ -166,14 +175,10 @@ func newPySparkEntrypointStrategy(execCtx *executionContext) entrypointStrategy resultURI: execCtx.s3aResultURI, returnResult: jobContext.ReturnResult, arguments: jobContext.Arguments, - bundleURI: execCtx.commandContext.BundleURI, } - - if jobContext.Parameters != nil && strings.TrimSpace(jobContext.Parameters.ScriptURI) != "" { - s.scriptURI = updateS3ToS3aURI(strings.TrimSpace(jobContext.Parameters.ScriptURI)) - return s + if queryURI := pysparkQueryURI(execCtx.commandContext, jobContext); queryURI != "" { + s.queryURI = queryURI } - if jobContext.Parameters != nil { s.entryPoint = jobContext.Parameters.EntryPoint s.bundleVersion = jobContext.Parameters.BundleVersion diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index 2ec0bf14..a4de9856 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -89,22 +89,22 @@ var ( ) var ( - ErrJobCanceled = fmt.Errorf("job was canceled before completion") - ErrJobSubmission = fmt.Errorf("failed to submit Spark application to Kubernetes cluster") - ErrKubeConfig = fmt.Errorf("failed to configure Kubernetes client: ensure EKS cluster access is properly configured") - ErrApplicationSpec = fmt.Errorf("failed to load or parse SparkApplication template") - ErrSparkApplicationFile = fmt.Errorf("failed to read SparkApplication application template file: check file path and permissions") - ErrMissingEntryPoint = fmt.Errorf("entry_point is required for .jar wrapper_uri: set parameters.entry_point to the fully-qualified main class") - ErrMissingBundleVersion = fmt.Errorf("parameters.bundle_version is required when the command sets bundle_uri") - ErrMissingBundleEntry = fmt.Errorf("parameters.entry_point is required when the command sets bundle_uri") - ErrInvalidScriptURI = fmt.Errorf("parameters.script_uri must be an s3:// or s3a:// .py object under the command bundle_uri prefix") - ErrConflictingPySparkSource = fmt.Errorf("parameters.script_uri and parameters.bundle_version are mutually exclusive") + ErrJobCanceled = fmt.Errorf("job was canceled before completion") + ErrJobSubmission = fmt.Errorf("failed to submit Spark application to Kubernetes cluster") + ErrKubeConfig = fmt.Errorf("failed to configure Kubernetes client: ensure EKS cluster access is properly configured") + ErrApplicationSpec = fmt.Errorf("failed to load or parse SparkApplication template") + ErrSparkApplicationFile = fmt.Errorf("failed to read SparkApplication application template file: check file path and permissions") + ErrMissingEntryPoint = fmt.Errorf("entry_point is required for .jar wrapper_uri: set parameters.entry_point to the fully-qualified main class") + ErrMissingBundleVersion = fmt.Errorf("parameters.bundle_version is required when the command sets bundle_uri") + ErrMissingBundleEntry = fmt.Errorf("parameters.entry_point is required when the command sets bundle_uri") + ErrInvalidScriptURI = fmt.Errorf("parameters.script_uri must be an s3:// or s3a:// .py object under the command bundle_uri prefix") + ErrConflictingPySparkSource = fmt.Errorf("parameters.script_uri and parameters.bundle_version are mutually exclusive") ) type commandContext struct { - JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` - WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_uri,omitempty"` - BundleURI string `yaml:"bundle_uri,omitempty" json:"bundle_uri,omitempty"` + JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` + WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_uri,omitempty"` + BundleURI string `yaml:"bundle_uri,omitempty" json:"bundle_uri,omitempty"` Image string `yaml:"image,omitempty" json:"image,omitempty"` EventLogURI string `yaml:"event_log_uri,omitempty" json:"event_log_uri,omitempty"` Properties map[string]string `yaml:"properties,omitempty" json:"properties,omitempty"` @@ -397,15 +397,19 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. // Set URIs and App Name execCtx.appName = fmt.Sprintf("%s-%s", applicationPrefix, j.ID) - execCtx.queryURI = fmt.Sprintf("%s/%s/%s/%s", s.JobsURI, j.ID, queriesPath, queryFileName) execCtx.resultURI = fmt.Sprintf("%s/%s/%s", s.JobsURI, j.ID, resultsPath) - execCtx.s3aQueryURI = updateS3ToS3aURI(execCtx.queryURI) execCtx.s3aResultURI = updateS3ToS3aURI(execCtx.resultURI) execCtx.logURI = fmt.Sprintf("%s/%s/%s", s.JobsURI, j.ID, logsPath) - // Upload query to S3 - if err := uploadFileToS3(ctx, execCtx.awsConfig, execCtx.queryURI, execCtx.jobContext.Query); err != nil { - return nil, fmt.Errorf("failed to upload query to S3: %w", err) + if queryURI := pysparkQueryURI(s, execCtx.jobContext); queryURI != "" { + execCtx.queryURI = queryURI + execCtx.s3aQueryURI = queryURI + } else { + execCtx.queryURI = fmt.Sprintf("%s/%s/%s/%s", s.JobsURI, j.ID, queriesPath, queryFileName) + execCtx.s3aQueryURI = updateS3ToS3aURI(execCtx.queryURI) + if err := uploadFileToS3(ctx, execCtx.awsConfig, execCtx.queryURI, execCtx.jobContext.Query); err != nil { + return nil, fmt.Errorf("failed to upload query to S3: %w", err) + } } // create empty log s3 directory to avoid spark event log dir errors @@ -850,7 +854,7 @@ func applySparkOperatorConfig(execCtx *executionContext) error { // Required Spark SQL extensions must win over job/cluster properties, so this // merge runs last: a caller can otherwise submit spark.sql.extensions="" and - // disable the Ranger authorization extension entirely. + // disable all the extensions. if clusterContext.RequiredSparkSQLExtensions != "" { existingExtensions := sparkApp.Spec.SparkConf[sparkSqlExtensions] if existingExtensions == "" { From 3799f2281cfc92cb1f1c6794494d57cecbdcf77e Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Tue, 8 Sep 2026 15:23:25 +0530 Subject: [PATCH 5/7] Changed approach --- .../pkg/object/command/sparkeks/entrypoint.go | 108 +++--------------- .../pkg/object/command/sparkeks/sparkeks.go | 33 ++---- 2 files changed, 26 insertions(+), 115 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index 3fea9048..8b7f73f3 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -76,70 +76,16 @@ func (s sqlWrapperEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) return nil } -type pysparkEntrypointStrategy struct { - appName string - queryURI string - user string - resultURI string - returnResult bool - arguments []string - bundleVersion string // job passes parameters.bundle_version - entryPoint string // job passes parameters.entry_point (path inside the zip) -} - -func (s pysparkEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) error { - if strings.HasSuffix(strings.ToLower(s.queryURI), ".py") { - spec.Arguments = buildArguments(s.arguments, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) - return nil - } - - if strings.TrimSpace(s.bundleVersion) == "" { - return ErrMissingBundleVersion - } - entryPoint := strings.TrimSpace(s.entryPoint) - if entryPoint == "" { - return ErrMissingBundleEntry - } - - extra := append([]string{entryPoint}, s.arguments...) - spec.Arguments = buildArguments(extra, s.appName, s.queryURI, s.user, s.resultURI, s.returnResult) - return nil -} - -func pysparkQueryURI(cmd *commandContext, jobCtx *jobContext) string { - if cmd == nil || strings.TrimSpace(cmd.BundleURI) == "" || jobCtx == nil || jobCtx.Parameters == nil { +// scriptQueryURI is parameters.script_uri as s3a for query_uri. Empty means upload query.sql. +func scriptQueryURI(jobCtx *jobContext) string { + if jobCtx == nil || jobCtx.Parameters == nil { return "" } - if script := strings.TrimSpace(jobCtx.Parameters.ScriptURI); script != "" { - return updateS3ToS3aURI(script) - } - if ver := strings.TrimSpace(jobCtx.Parameters.BundleVersion); ver != "" { - return updateS3ToS3aURI(strings.TrimRight(cmd.BundleURI, "/") + "/" + ver + ".zip") - } - return "" -} - -func validateScriptURI(scriptURI, allowedPrefix string) error { - scriptURI = strings.TrimSpace(scriptURI) - if scriptURI == "" { - return ErrInvalidScriptURI - } - if !strings.HasPrefix(scriptURI, s3Prefix) && !strings.HasPrefix(scriptURI, s3aPrefix) { - return ErrInvalidScriptURI - } - if strings.Contains(scriptURI, "..") { - return ErrInvalidScriptURI - } - if !strings.HasSuffix(strings.ToLower(scriptURI), ".py") { - return ErrInvalidScriptURI - } - - normalized := updateS3ToS3aURI(scriptURI) - prefix := updateS3ToS3aURI(strings.TrimRight(strings.TrimSpace(allowedPrefix), "/")) + "/" - if !strings.HasPrefix(normalized, prefix) { - return ErrInvalidScriptURI + script := strings.TrimSpace(jobCtx.Parameters.ScriptURI) + if script == "" { + return "" } - return nil + return updateS3ToS3aURI(script) } // entrypointFactory builds the entrypoint strategy for a job from its execution context. @@ -153,10 +99,6 @@ var entrypointStrategiesByExt = map[string]entrypointFactory{ var defaultEntrypointFactory entrypointFactory = newSQLWrapperEntrypointStrategy func newEntrypointStrategy(execCtx *executionContext) entrypointStrategy { - if execCtx.commandContext.BundleURI != "" { - return newPySparkEntrypointStrategy(execCtx) - } - ext := strings.ToLower(path.Ext(execCtx.commandContext.WrapperURI)) factory, ok := entrypointStrategiesByExt[ext] if !ok { @@ -165,27 +107,6 @@ func newEntrypointStrategy(execCtx *executionContext) entrypointStrategy { return factory(execCtx) } -func newPySparkEntrypointStrategy(execCtx *executionContext) entrypointStrategy { - jobContext := execCtx.jobContext - - s := pysparkEntrypointStrategy{ - appName: execCtx.appName, - queryURI: execCtx.s3aQueryURI, - user: execCtx.job.User, - resultURI: execCtx.s3aResultURI, - returnResult: jobContext.ReturnResult, - arguments: jobContext.Arguments, - } - if queryURI := pysparkQueryURI(execCtx.commandContext, jobContext); queryURI != "" { - s.queryURI = queryURI - } - if jobContext.Parameters != nil { - s.entryPoint = jobContext.Parameters.EntryPoint - s.bundleVersion = jobContext.Parameters.BundleVersion - } - return s -} - func newJarEntrypointStrategy(execCtx *executionContext) entrypointStrategy { jobContext := execCtx.jobContext @@ -209,13 +130,22 @@ func newJarEntrypointStrategy(execCtx *executionContext) entrypointStrategy { func newSQLWrapperEntrypointStrategy(execCtx *executionContext) entrypointStrategy { jobContext := execCtx.jobContext - + queryURI := execCtx.s3aQueryURI + if u := scriptQueryURI(jobContext); u != "" { + queryURI = u + } + extra := jobContext.Arguments + if jobContext.Parameters != nil { + if ep := strings.TrimSpace(jobContext.Parameters.EntryPoint); ep != "" { + extra = append([]string{ep}, extra...) + } + } return sqlWrapperEntrypointStrategy{ appName: execCtx.appName, - queryURI: execCtx.s3aQueryURI, + queryURI: queryURI, user: execCtx.job.User, resultURI: execCtx.s3aResultURI, returnResult: jobContext.ReturnResult, - arguments: jobContext.Arguments, + arguments: extra, } } diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index a4de9856..0da4f357 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -89,22 +89,17 @@ var ( ) var ( - ErrJobCanceled = fmt.Errorf("job was canceled before completion") - ErrJobSubmission = fmt.Errorf("failed to submit Spark application to Kubernetes cluster") - ErrKubeConfig = fmt.Errorf("failed to configure Kubernetes client: ensure EKS cluster access is properly configured") - ErrApplicationSpec = fmt.Errorf("failed to load or parse SparkApplication template") - ErrSparkApplicationFile = fmt.Errorf("failed to read SparkApplication application template file: check file path and permissions") - ErrMissingEntryPoint = fmt.Errorf("entry_point is required for .jar wrapper_uri: set parameters.entry_point to the fully-qualified main class") - ErrMissingBundleVersion = fmt.Errorf("parameters.bundle_version is required when the command sets bundle_uri") - ErrMissingBundleEntry = fmt.Errorf("parameters.entry_point is required when the command sets bundle_uri") - ErrInvalidScriptURI = fmt.Errorf("parameters.script_uri must be an s3:// or s3a:// .py object under the command bundle_uri prefix") - ErrConflictingPySparkSource = fmt.Errorf("parameters.script_uri and parameters.bundle_version are mutually exclusive") + ErrJobCanceled = fmt.Errorf("job was canceled before completion") + ErrJobSubmission = fmt.Errorf("failed to submit Spark application to Kubernetes cluster") + ErrKubeConfig = fmt.Errorf("failed to configure Kubernetes client: ensure EKS cluster access is properly configured") + ErrApplicationSpec = fmt.Errorf("failed to load or parse SparkApplication template") + ErrSparkApplicationFile = fmt.Errorf("failed to read SparkApplication application template file: check file path and permissions") + ErrMissingEntryPoint = fmt.Errorf("entry_point is required for .jar wrapper_uri: set parameters.entry_point to the fully-qualified main class") ) type commandContext struct { JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_uri,omitempty"` - BundleURI string `yaml:"bundle_uri,omitempty" json:"bundle_uri,omitempty"` Image string `yaml:"image,omitempty" json:"image,omitempty"` EventLogURI string `yaml:"event_log_uri,omitempty" json:"event_log_uri,omitempty"` Properties map[string]string `yaml:"properties,omitempty" json:"properties,omitempty"` @@ -115,7 +110,6 @@ type jobParameters struct { Properties map[string]string `yaml:"properties,omitempty" json:"properties,omitempty"` EntryPoint string `yaml:"entry_point,omitempty" json:"entry_point,omitempty"` ApplicationType string `yaml:"application_type,omitempty" json:"application_type,omitempty"` - BundleVersion string `yaml:"bundle_version,omitempty" json:"bundle_version,omitempty"` ScriptURI string `yaml:"script_uri,omitempty" json:"script_uri,omitempty"` } @@ -348,19 +342,6 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. } execCtx.jobContext = jobContext - if s.BundleURI != "" && jobContext.Parameters != nil { - scriptURI := strings.TrimSpace(jobContext.Parameters.ScriptURI) - bundleVersion := strings.TrimSpace(jobContext.Parameters.BundleVersion) - if scriptURI != "" && bundleVersion != "" { - return nil, ErrConflictingPySparkSource - } - if scriptURI != "" { - if err := validateScriptURI(scriptURI, s.BundleURI); err != nil { - return nil, err - } - } - } - // Parse cluster context clusterContext := &clusterContext{} if c.Context != nil { @@ -401,7 +382,7 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. execCtx.s3aResultURI = updateS3ToS3aURI(execCtx.resultURI) execCtx.logURI = fmt.Sprintf("%s/%s/%s", s.JobsURI, j.ID, logsPath) - if queryURI := pysparkQueryURI(s, execCtx.jobContext); queryURI != "" { + if queryURI := scriptQueryURI(execCtx.jobContext); queryURI != "" { execCtx.queryURI = queryURI execCtx.s3aQueryURI = queryURI } else { From 3ff4b9ccd952c71ab2ea56b0683369a5fdbe4681 Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Tue, 8 Sep 2026 15:35:30 +0530 Subject: [PATCH 6/7] Changes --- .../pkg/object/command/sparkeks/entrypoint.go | 18 +----------------- .../pkg/object/command/sparkeks/sparkeks.go | 6 +++--- 2 files changed, 4 insertions(+), 20 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index 8b7f73f3..36201e3c 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -76,18 +76,6 @@ func (s sqlWrapperEntrypointStrategy) apply(spec *v1beta2.SparkApplicationSpec) return nil } -// scriptQueryURI is parameters.script_uri as s3a for query_uri. Empty means upload query.sql. -func scriptQueryURI(jobCtx *jobContext) string { - if jobCtx == nil || jobCtx.Parameters == nil { - return "" - } - script := strings.TrimSpace(jobCtx.Parameters.ScriptURI) - if script == "" { - return "" - } - return updateS3ToS3aURI(script) -} - // entrypointFactory builds the entrypoint strategy for a job from its execution context. type entrypointFactory func(execCtx *executionContext) entrypointStrategy @@ -130,10 +118,6 @@ func newJarEntrypointStrategy(execCtx *executionContext) entrypointStrategy { func newSQLWrapperEntrypointStrategy(execCtx *executionContext) entrypointStrategy { jobContext := execCtx.jobContext - queryURI := execCtx.s3aQueryURI - if u := scriptQueryURI(jobContext); u != "" { - queryURI = u - } extra := jobContext.Arguments if jobContext.Parameters != nil { if ep := strings.TrimSpace(jobContext.Parameters.EntryPoint); ep != "" { @@ -142,7 +126,7 @@ func newSQLWrapperEntrypointStrategy(execCtx *executionContext) entrypointStrate } return sqlWrapperEntrypointStrategy{ appName: execCtx.appName, - queryURI: queryURI, + queryURI: execCtx.s3aQueryURI, user: execCtx.job.User, resultURI: execCtx.s3aResultURI, returnResult: jobContext.ReturnResult, diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index 0da4f357..5a91829e 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -382,9 +382,9 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. execCtx.s3aResultURI = updateS3ToS3aURI(execCtx.resultURI) execCtx.logURI = fmt.Sprintf("%s/%s/%s", s.JobsURI, j.ID, logsPath) - if queryURI := scriptQueryURI(execCtx.jobContext); queryURI != "" { - execCtx.queryURI = queryURI - execCtx.s3aQueryURI = queryURI + if script := strings.TrimSpace(execCtx.jobContext.Parameters.ScriptURI); script != "" { + execCtx.queryURI = script + execCtx.s3aQueryURI = updateS3ToS3aURI(script) } else { execCtx.queryURI = fmt.Sprintf("%s/%s/%s/%s", s.JobsURI, j.ID, queriesPath, queryFileName) execCtx.s3aQueryURI = updateS3ToS3aURI(execCtx.queryURI) From 7a8cecd410b5fbb4e37b1752844e5f4c92358bdb Mon Sep 17 00:00:00 2001 From: Nitin Bhakar Date: Tue, 8 Sep 2026 21:26:31 +0530 Subject: [PATCH 7/7] Review changes --- .../pkg/object/command/sparkeks/entrypoint.go | 2 +- .../command/sparkeks/entrypoint_test.go | 48 +++++++++++++++++++ .../pkg/object/command/sparkeks/sparkeks.go | 22 ++++++--- 3 files changed, 65 insertions(+), 7 deletions(-) diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index 36201e3c..ff6577bd 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -121,7 +121,7 @@ func newSQLWrapperEntrypointStrategy(execCtx *executionContext) entrypointStrate extra := jobContext.Arguments if jobContext.Parameters != nil { if ep := strings.TrimSpace(jobContext.Parameters.EntryPoint); ep != "" { - extra = append([]string{ep}, extra...) + extra = append(extra, ep) } } return sqlWrapperEntrypointStrategy{ diff --git a/internal/pkg/object/command/sparkeks/entrypoint_test.go b/internal/pkg/object/command/sparkeks/entrypoint_test.go index 17db798a..097d8890 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint_test.go +++ b/internal/pkg/object/command/sparkeks/entrypoint_test.go @@ -177,6 +177,54 @@ func TestSQLWrapperEntrypointStrategy_Apply_NoEntryPointRequired(t *testing.T) { } } +func TestSQLWrapperEntrypointStrategy_ScriptURISkipsQuerySQL(t *testing.T) { + jobCtx := &jobContext{ + ReturnResult: true, + Arguments: []string{"--date", "2026-09-01"}, + Parameters: &jobParameters{ + ScriptURI: "s3://bucket/pyspark/abc123.zip", + EntryPoint: "src/job.py", + }, + } + execCtx := newTestExecutionContext("s3://bucket/wrapper.py", jobCtx, "alice", "result_uri") + + if assignQueryURI(execCtx, "s3://jobs", "job-id") { + t.Fatal("assignQueryURI() = true, want false (do not upload query.sql)") + } + if execCtx.s3aQueryURI != "s3a://bucket/pyspark/abc123.zip" { + t.Fatalf("s3aQueryURI = %q, want script_uri as s3a", execCtx.s3aQueryURI) + } + + spec := &v1beta2.SparkApplicationSpec{} + if err := newEntrypointStrategy(execCtx).apply(spec); err != nil { + t.Fatalf("apply() returned unexpected error: %v", err) + } + + expected := []string{ + "spark-sql-job-test", + "s3a://bucket/pyspark/abc123.zip", + "alice", + "result_uri", + "--date", + "2026-09-01", + "src/job.py", + } + if !reflect.DeepEqual(spec.Arguments, expected) { + t.Errorf("spec.Arguments = %v, want %v", spec.Arguments, expected) + } +} + +func TestAssignQueryURI_SQLUploadsQuerySQL(t *testing.T) { + execCtx := newTestExecutionContext("s3://bucket/wrapper.py", &jobContext{}, "alice", "result_uri") + if !assignQueryURI(execCtx, "s3://jobs", "job-id") { + t.Fatal("assignQueryURI() = false, want true (upload query.sql)") + } + want := "s3a://jobs/job-id/queries/query.sql" + if execCtx.s3aQueryURI != want { + t.Errorf("s3aQueryURI = %q, want %q", execCtx.s3aQueryURI, want) + } +} + func newTestExecutionContext(wrapperURI string, jobCtx *jobContext, user string, resultURI string) *executionContext { return &executionContext{ job: &job.Job{ diff --git a/internal/pkg/object/command/sparkeks/sparkeks.go b/internal/pkg/object/command/sparkeks/sparkeks.go index 5a91829e..d4dd787b 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -382,12 +382,7 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. execCtx.s3aResultURI = updateS3ToS3aURI(execCtx.resultURI) execCtx.logURI = fmt.Sprintf("%s/%s/%s", s.JobsURI, j.ID, logsPath) - if script := strings.TrimSpace(execCtx.jobContext.Parameters.ScriptURI); script != "" { - execCtx.queryURI = script - execCtx.s3aQueryURI = updateS3ToS3aURI(script) - } else { - execCtx.queryURI = fmt.Sprintf("%s/%s/%s/%s", s.JobsURI, j.ID, queriesPath, queryFileName) - execCtx.s3aQueryURI = updateS3ToS3aURI(execCtx.queryURI) + if assignQueryURI(execCtx, s.JobsURI, j.ID) { if err := uploadFileToS3(ctx, execCtx.awsConfig, execCtx.queryURI, execCtx.jobContext.Query); err != nil { return nil, fmt.Errorf("failed to upload query to S3: %w", err) } @@ -401,6 +396,21 @@ func buildExecutionContextAndURI(ctx context.Context, r *plugin.Runtime, j *job. return execCtx, nil } +// assignQueryURI sets queryURI from parameters.script_uri, or the uploaded query.sql path. +// Returns true when the caller should upload query.sql. +func assignQueryURI(execCtx *executionContext, jobsURI, jobID string) bool { + if execCtx.jobContext != nil && execCtx.jobContext.Parameters != nil { + if script := strings.TrimSpace(execCtx.jobContext.Parameters.ScriptURI); script != "" { + execCtx.queryURI = script + execCtx.s3aQueryURI = updateS3ToS3aURI(script) + return false + } + } + execCtx.queryURI = fmt.Sprintf("%s/%s/%s/%s", jobsURI, jobID, queriesPath, queryFileName) + execCtx.s3aQueryURI = updateS3ToS3aURI(execCtx.queryURI) + return true +} + // submitSparkApp creates clients, generates the spec, and submits it to Kubernetes. func (e *executionContext) submitSparkApp(ctx context.Context) error { // Create Kubernetes and Spark Operator clients