diff --git a/internal/pkg/object/command/sparkeks/entrypoint.go b/internal/pkg/object/command/sparkeks/entrypoint.go index f5aebb7..ff6577b 100644 --- a/internal/pkg/object/command/sparkeks/entrypoint.go +++ b/internal/pkg/object/command/sparkeks/entrypoint.go @@ -118,13 +118,18 @@ func newJarEntrypointStrategy(execCtx *executionContext) entrypointStrategy { func newSQLWrapperEntrypointStrategy(execCtx *executionContext) entrypointStrategy { jobContext := execCtx.jobContext - + extra := jobContext.Arguments + if jobContext.Parameters != nil { + if ep := strings.TrimSpace(jobContext.Parameters.EntryPoint); ep != "" { + extra = append(extra, ep) + } + } return sqlWrapperEntrypointStrategy{ appName: execCtx.appName, queryURI: execCtx.s3aQueryURI, user: execCtx.job.User, resultURI: execCtx.s3aResultURI, returnResult: jobContext.ReturnResult, - arguments: jobContext.Arguments, + arguments: extra, } } diff --git a/internal/pkg/object/command/sparkeks/entrypoint_test.go b/internal/pkg/object/command/sparkeks/entrypoint_test.go index 17db798..097d889 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 6df2f2e..d4dd787 100644 --- a/internal/pkg/object/command/sparkeks/sparkeks.go +++ b/internal/pkg/object/command/sparkeks/sparkeks.go @@ -100,6 +100,7 @@ var ( type commandContext struct { JobsURI string `yaml:"jobs_uri,omitempty" json:"jobs_uri,omitempty"` WrapperURI string `yaml:"wrapper_uri,omitempty" json:"wrapper_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 +110,7 @@ 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"` + ScriptURI string `yaml:"script_uri,omitempty" json:"script_uri,omitempty"` } type jobContext struct { @@ -376,15 +378,14 @@ 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 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) + } } // create empty log s3 directory to avoid spark event log dir errors @@ -395,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 @@ -727,6 +743,17 @@ func updateKubeConfig(ctx context.Context, execCtx *executionContext) (string, e return kubeconfigPath, nil } +// 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 + } + 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 +793,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 +804,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 +843,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 all the extensions. + 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 }