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
9 changes: 7 additions & 2 deletions internal/pkg/object/command/sparkeks/entrypoint.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
}
}
48 changes: 48 additions & 0 deletions internal/pkg/object/command/sparkeks/entrypoint_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down
108 changes: 66 additions & 42 deletions internal/pkg/object/command/sparkeks/sparkeks.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand All @@ -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 {
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand All @@ -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 != "" {
Expand Down Expand Up @@ -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
}
Expand Down
Loading