This is an automated email from the ASF dual-hosted git repository.

Abacn pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git


The following commit(s) were added to refs/heads/master by this push:
     new 830d7b63ab5 [Spark] Make cancel() cancel the Spark jobs and stop only 
a session the runner created (#40103)
830d7b63ab5 is described below

commit 830d7b63ab5d7dac48485dd5dd0f1e0b7d135b0b
Author: Tobias Kaymak <[email protected]>
AuthorDate: Mon Sep 14 17:13:18 2026 +0200

    [Spark] Make cancel() cancel the Spark jobs and stop only a session the 
runner created (#40103)
    
    * [Spark] Make cancel() cancel the Spark jobs and stop only a session the 
runner created
    
    Fixes #40101.
---
 ...eam_PostCommit_Java_ValidatesRunner_Spark4.json |  2 +-
 ...a_ValidatesRunner_SparkStructuredStreaming.json |  3 +-
 .../translation/StreamingEvaluationContext.java    | 11 ++--
 .../SparkStructuredStreamingPipelineResult.java    | 74 ++++++++++++----------
 .../SparkStructuredStreamingRunner.java            | 48 +++++++++-----
 .../translation/EvaluationContext.java             | 20 ++++--
 .../translation/SparkSessionFactory.java           | 48 ++++++++++++--
 ...SparkStructuredStreamingPipelineResultTest.java | 66 +++++++++++++++++++
 .../StructuredStreamingPipelineStateTest.java      | 72 +++++++++++++++++++++
 9 files changed, 278 insertions(+), 66 deletions(-)

diff --git 
a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_Spark4.json 
b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_Spark4.json
index e3d6056a5de..b2683333323 100644
--- a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_Spark4.json
+++ b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_Spark4.json
@@ -1,4 +1,4 @@
 {
   "comment": "Modify this file in a trivial way to cause this test suite to 
run",
-  "modification": 1
+  "modification": 2
 }
diff --git 
a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
 
b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
index cad8d98b8ea..373c31ff234 100644
--- 
a/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
+++ 
b/.github/trigger_files/beam_PostCommit_Java_ValidatesRunner_SparkStructuredStreaming.json
@@ -8,5 +8,6 @@
   "https://github.com/apache/beam/pull/34080": "noting that PR #34080 should 
run this test",
   "https://github.com/apache/beam/pull/34155": "noting that PR #34155 should 
run this test",
   "https://github.com/apache/beam/pull/35159": "moving WindowedValue and 
making an interface",
-  "https://github.com/apache/beam/pull/39793": "noting that PR #39793 should 
run this test"
+  "https://github.com/apache/beam/pull/39793": "noting that PR #39793 should 
run this test",
+  "https://github.com/apache/beam/pull/40103": "noting that PR #40103 should 
run this test"
 }
diff --git 
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/StreamingEvaluationContext.java
 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/StreamingEvaluationContext.java
index ff7839bddb2..2ca54a1e948 100644
--- 
a/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/StreamingEvaluationContext.java
+++ 
b/runners/spark/4/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/StreamingEvaluationContext.java
@@ -56,10 +56,9 @@ public class StreamingEvaluationContext extends 
EvaluationContext {
 
   private final SparkStructuredStreamingPipelineOptions options;
 
-  // Guards queries and stopped.
+  // Guards queries and the stopped flag.
   private final Object lock = new Object();
   private final List<StreamingQuery> queries = new ArrayList<>();
-  private boolean stopped = false;
 
   StreamingEvaluationContext(
       Collection<? extends NamedDataset<?>> leaves,
@@ -92,7 +91,7 @@ public class StreamingEvaluationContext extends 
EvaluationContext {
           continue;
         }
         synchronized (lock) {
-          if (stopped) {
+          if (isStopped()) {
             break;
           }
         }
@@ -105,7 +104,7 @@ public class StreamingEvaluationContext extends 
EvaluationContext {
         boolean alreadyStopped;
         synchronized (lock) {
           queries.add(query);
-          alreadyStopped = stopped;
+          alreadyStopped = isStopped();
         }
         if (alreadyStopped) {
           stopQuery(query);
@@ -134,10 +133,10 @@ public class StreamingEvaluationContext extends 
EvaluationContext {
   public void stop() {
     List<StreamingQuery> toStop;
     synchronized (lock) {
-      if (stopped) {
+      if (isStopped()) {
         return;
       }
-      stopped = true;
+      super.stop();
       toStop = new ArrayList<>(queries);
     }
     for (StreamingQuery query : toStop) {
diff --git 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResult.java
 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResult.java
index b592b6fb742..c483bbcf3cf 100644
--- 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResult.java
+++ 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResult.java
@@ -25,6 +25,7 @@ import java.util.concurrent.ExecutionException;
 import java.util.concurrent.Future;
 import java.util.concurrent.TimeUnit;
 import java.util.concurrent.TimeoutException;
+import java.util.concurrent.atomic.AtomicBoolean;
 import java.util.function.Supplier;
 import 
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
 import 
org.apache.beam.runners.spark.structuredstreaming.translation.EvaluationContext;
@@ -35,25 +36,37 @@ import org.apache.beam.sdk.util.UserCodeException;
 import org.apache.spark.SparkException;
 import org.checkerframework.checker.nullness.qual.Nullable;
 import org.joda.time.Duration;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
 
+/**
+ * Result of a pipeline submitted to the {@link 
SparkStructuredStreamingRunner}. The pipeline runs
+ * asynchronously on a dedicated thread.
+ */
 public class SparkStructuredStreamingPipelineResult implements PipelineResult {
 
+  private static final Logger LOG =
+      LoggerFactory.getLogger(SparkStructuredStreamingPipelineResult.class);
+
   private final Future<?> pipelineExecution;
   // Supplies the context of the translated pipeline, null until translation 
has completed.
   private final Supplier<? extends @Nullable EvaluationContext> 
evaluationContext;
   private final MetricsAccumulator metrics;
-  private final @Nullable Runnable onTerminalState;
-  private PipelineResult.State state;
+  private final AtomicBoolean cancelRequested;
+  private final Runnable cancelSparkJobs;
+  private volatile PipelineResult.State state;
 
   SparkStructuredStreamingPipelineResult(
       Future<?> pipelineExecution,
       Supplier<? extends @Nullable EvaluationContext> evaluationContext,
       MetricsAccumulator metrics,
-      final @Nullable Runnable onTerminalState) {
+      AtomicBoolean cancelRequested,
+      Runnable cancelSparkJobs) {
     this.pipelineExecution = pipelineExecution;
     this.evaluationContext = evaluationContext;
     this.metrics = metrics;
-    this.onTerminalState = onTerminalState;
+    this.cancelRequested = cancelRequested;
+    this.cancelSparkJobs = cancelSparkJobs;
     // pipelineExecution is expected to have started executing eagerly.
     this.state = State.RUNNING;
   }
@@ -77,13 +90,6 @@ public class SparkStructuredStreamingPipelineResult 
implements PipelineResult {
         : new Pipeline.PipelineExecutionException(firstNonNull(next, 
exception));
   }
 
-  private State awaitTermination(Duration duration)
-      throws TimeoutException, ExecutionException, InterruptedException {
-    pipelineExecution.get(duration.getMillis(), TimeUnit.MILLISECONDS);
-    // Throws an exception if the job is not finished successfully in the 
given time.
-    return PipelineResult.State.DONE;
-  }
-
   @Override
   public PipelineResult.State getState() {
     return state;
@@ -94,18 +100,31 @@ public class SparkStructuredStreamingPipelineResult 
implements PipelineResult {
     return waitUntilFinish(Duration.millis(Long.MAX_VALUE));
   }
 
+  /**
+   * Waits up to {@code duration} for the execution thread. A pipeline that 
ends after {@link
+   * #cancel()} is CANCELLED, any other failure is rethrown and the pipeline 
is FAILED.
+   */
   @Override
   public State waitUntilFinish(final Duration duration) {
     try {
-      State finishState = awaitTermination(duration);
-      offerNewState(finishState);
+      pipelineExecution.get(duration.getMillis(), TimeUnit.MILLISECONDS);
+      state = cancelRequested.get() ? State.CANCELLED : State.DONE;
     } catch (final TimeoutException e) {
       // ignore.
     } catch (final ExecutionException e) {
-      offerNewState(PipelineResult.State.FAILED);
+      if (cancelRequested.get()) {
+        LOG.warn("Pipeline execution failed after cancel", e.getCause());
+        state = State.CANCELLED;
+        return state;
+      }
+      state = State.FAILED;
       throw unwrapCause(firstNonNull(e.getCause(), e));
+    } catch (final InterruptedException e) {
+      Thread.currentThread().interrupt();
+      state = State.FAILED;
+      throw unwrapCause(e);
     } catch (final Exception e) {
-      offerNewState(PipelineResult.State.FAILED);
+      state = State.FAILED;
       throw unwrapCause(e);
     }
 
@@ -117,26 +136,17 @@ public class SparkStructuredStreamingPipelineResult 
implements PipelineResult {
     return asAttemptedOnlyMetricResults(metrics.value());
   }
 
+  /** Requests cancellation of the pipeline and returns immediately. */
   @Override
   public PipelineResult.State cancel() throws IOException {
-    EvaluationContext ctx = evaluationContext.get();
-    if (ctx != null) {
-      ctx.stop();
-    }
-    pipelineExecution.cancel(true);
-    offerNewState(PipelineResult.State.CANCELLED);
-    return state;
-  }
-
-  private void offerNewState(State newState) {
-    State oldState = this.state;
-    this.state = newState;
-    if (!oldState.isTerminal() && newState.isTerminal() && onTerminalState != 
null) {
-      try {
-        onTerminalState.run();
-      } catch (Exception e) {
-        throw unwrapCause(e);
+    if (!state.isTerminal() && cancelRequested.compareAndSet(false, true)) {
+      EvaluationContext ctx = evaluationContext.get();
+      if (ctx != null) {
+        ctx.stop();
       }
+      cancelSparkJobs.run();
+      state = State.CANCELLED;
     }
+    return state;
   }
 }
diff --git 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingRunner.java
 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingRunner.java
index f78026847fa..0f5ddffae34 100644
--- 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingRunner.java
+++ 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingRunner.java
@@ -17,12 +17,13 @@
  */
 package org.apache.beam.runners.spark.structuredstreaming;
 
+import java.util.UUID;
 import java.util.concurrent.ExecutorService;
 import java.util.concurrent.Executors;
 import java.util.concurrent.Future;
 import java.util.concurrent.ThreadFactory;
+import java.util.concurrent.atomic.AtomicBoolean;
 import java.util.concurrent.atomic.AtomicReference;
-import javax.annotation.Nullable;
 import org.apache.beam.runners.core.metrics.MetricsPusher;
 import org.apache.beam.runners.core.metrics.NoOpMetricsSink;
 import 
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
@@ -42,6 +43,7 @@ import org.apache.beam.sdk.options.PipelineOptionsValidator;
 import org.apache.beam.sdk.util.construction.SplittableParDo;
 import org.apache.beam.sdk.util.construction.graph.ProjectionPushdownOptimizer;
 import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ThreadFactoryBuilder;
+import org.apache.spark.SparkContext;
 import org.apache.spark.SparkEnv$;
 import org.apache.spark.metrics.MetricsSystem;
 import org.apache.spark.sql.SparkSession;
@@ -145,27 +147,47 @@ public final class SparkStructuredStreamingRunner
 
     PipelineTranslator.detectStreamingMode(pipeline, options);
 
-    final SparkSession sparkSession = 
SparkSessionFactory.getOrCreateSession(options);
+    final boolean releaseSession = !options.getUseActiveSparkSession();
+    final SparkSession sparkSession = SparkSessionFactory.acquire(options);
+    final SparkContext sc = sparkSession.sparkContext();
     final MetricsAccumulator metrics = 
MetricsAccumulator.getInstance(sparkSession);
 
-    // Set once the pipeline is translated, so the result can stop an ongoing 
(streaming)
-    // evaluation on cancel. Remains null until translation completes.
+    // Null until translation completes.
     final AtomicReference<EvaluationContext> ctxRef = new AtomicReference<>();
+    final AtomicBoolean cancelRequested = new AtomicBoolean(false);
+
+    final String jobName = options.getJobName();
+    final String jobGroupId = "beam-" + jobName + "-" + UUID.randomUUID();
+    final Runnable cancelSparkJobs =
+        () -> {
+          try {
+            sc.cancelJobGroup(jobGroupId);
+          } catch (IllegalStateException e) {
+            // Context stopped concurrently.
+          }
+        };
 
     final Future<?> submissionFuture =
         runAsync(
             () -> {
-              EvaluationContext ctx = translatePipeline(sparkSession, 
pipeline);
-              ctxRef.set(ctx);
-              ctx.evaluate();
+              try {
+                // Interrupts running tasks on cancel, as Spark's 
StreamExecution does.
+                sc.setJobGroup(jobGroupId, "Beam " + jobName, true);
+                EvaluationContext ctx = translatePipeline(sparkSession, 
pipeline);
+                ctxRef.set(ctx);
+                if (!cancelRequested.get()) {
+                  ctx.evaluate();
+                }
+              } finally {
+                if (releaseSession) {
+                  SparkSessionFactory.release(sparkSession);
+                }
+              }
             });
 
     final SparkStructuredStreamingPipelineResult result =
         new SparkStructuredStreamingPipelineResult(
-            submissionFuture,
-            ctxRef::get,
-            metrics,
-            sparkStopFn(sparkSession, options.getUseActiveSparkSession()));
+            submissionFuture, ctxRef::get, metrics, cancelRequested, 
cancelSparkJobs);
 
     if (options.getEnableSparkMetricSinks()) {
       registerMetricsSource(options.getAppName(), metrics);
@@ -228,8 +250,4 @@ public final class SparkStructuredStreamingRunner
     execService.shutdown();
     return future;
   }
-
-  private static @Nullable Runnable sparkStopFn(SparkSession session, boolean 
isProvided) {
-    return !isProvided ? () -> session.stop() : null;
-  }
 }
diff --git 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/EvaluationContext.java
 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/EvaluationContext.java
index 0e677051fb6..fb559ab8d9c 100644
--- 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/EvaluationContext.java
+++ 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/EvaluationContext.java
@@ -52,6 +52,7 @@ public class EvaluationContext {
 
   private final Collection<? extends NamedDataset<?>> leaves;
   private final SparkSession session;
+  private volatile boolean stopped = false;
 
   protected EvaluationContext(Collection<? extends NamedDataset<?>> leaves, 
SparkSession session) {
     this.leaves = leaves;
@@ -63,9 +64,13 @@ public class EvaluationContext {
     return leaves;
   }
 
-  /** Trigger evaluation of all leaf datasets. */
+  /** Trigger evaluation of all leaf datasets. Returns early once {@link 
#stop()} was called. */
   public void evaluate() {
     for (NamedDataset<?> ds : leaves) {
+      if (stopped) {
+        LOG.info("Evaluation stopped, skipping remaining datasets");
+        return;
+      }
       final Dataset<?> dataset = ds.dataset();
       if (dataset == null) {
         continue;
@@ -119,11 +124,16 @@ public class EvaluationContext {
   }
 
   /**
-   * Stops any ongoing streaming execution triggered by this context.
-   *
-   * <p>This is a no-op for batch pipelines.
+   * Stops the evaluation after the current leaf dataset. Streaming contexts 
override this to stop
+   * their queries.
    */
-  public void stop() {}
+  public void stop() {
+    stopped = true;
+  }
+
+  protected boolean isStopped() {
+    return stopped;
+  }
 
   public SparkSession getSparkSession() {
     return session;
diff --git 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/SparkSessionFactory.java
 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/SparkSessionFactory.java
index 148188bb15a..4222857f4f3 100644
--- 
a/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/SparkSessionFactory.java
+++ 
b/runners/spark/src/main/java/org/apache/beam/runners/spark/structuredstreaming/translation/SparkSessionFactory.java
@@ -28,6 +28,7 @@ import java.util.ArrayList;
 import java.util.Collection;
 import java.util.Collections;
 import java.util.HashMap;
+import java.util.Map;
 import javax.annotation.Nullable;
 import org.apache.beam.repackaged.core.org.apache.commons.lang3.ArrayUtils;
 import org.apache.beam.runners.core.construction.SerializablePipelineOptions;
@@ -90,6 +91,7 @@ import org.apache.spark.sql.catalyst.InternalRow;
 import 
org.apache.spark.sql.execution.datasources.v2.DataWritingSparkTaskResult;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
+import scala.Option;
 
 public class SparkSessionFactory {
 
@@ -113,15 +115,49 @@ public class SparkSessionFactory {
           "/com.esotericsoftware/kryo-shaded",
           "/com/esotericsoftware/kryo-shaded");
 
-  /**
-   * Gets active {@link SparkSession} or creates one using {@link
-   * SparkStructuredStreamingPipelineOptions}.
-   */
-  public static SparkSession 
getOrCreateSession(SparkStructuredStreamingPipelineOptions options) {
+  // Builder.getOrCreate adopts an existing session without applying the 
pipeline's configuration.
+  // A pipeline must not stop a session it did not create, and the next 
pipeline needs the
+  // previous one stopped to get its own configuration, so sessions created 
here are counted.
+  private static final Map<SparkSession, Integer> OWNED_SESSIONS = new 
HashMap<>();
+
+  /** Returns the {@link SparkSession} for a pipeline, paired with {@link 
#release}. */
+  public static synchronized SparkSession 
acquire(SparkStructuredStreamingPipelineOptions options) {
     if (options.getUseActiveSparkSession()) {
       return SparkSession.active();
     }
-    return sessionBuilder(options.getSparkMaster(), options).getOrCreate();
+    boolean noUsableSession =
+        !isUsable(SparkSession.getActiveSession()) && 
!isUsable(SparkSession.getDefaultSession());
+    SparkSession session = sessionBuilder(options.getSparkMaster(), 
options).getOrCreate();
+    Integer count = OWNED_SESSIONS.get(session);
+    if (count != null) {
+      OWNED_SESSIONS.put(session, count + 1);
+      LOG.info("Pipeline options will not be applied to the shared 
SparkSession");
+    } else if (noUsableSession) {
+      OWNED_SESSIONS.put(session, 1);
+    }
+    return session;
+  }
+
+  /**
+   * Releases a session from {@link #acquire} and stops it when no longer 
used. The stop runs under
+   * the lock, a pipeline starting meanwhile creates a new session.
+   */
+  public static synchronized void release(SparkSession session) {
+    Integer count = OWNED_SESSIONS.get(session);
+    if (count == null) {
+      return;
+    }
+    if (count > 1) {
+      OWNED_SESSIONS.put(session, count - 1);
+      return;
+    }
+    OWNED_SESSIONS.remove(session);
+    LOG.info("Stopping SparkSession created by the runner");
+    session.stop();
+  }
+
+  private static boolean isUsable(Option<SparkSession> session) {
+    return session.isDefined() && !session.get().sparkContext().isStopped();
   }
 
   /** Creates Spark session builder with some optimizations for local mode, 
e.g. in tests. */
diff --git 
a/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResultTest.java
 
b/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResultTest.java
new file mode 100644
index 00000000000..4cd4858f000
--- /dev/null
+++ 
b/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/SparkStructuredStreamingPipelineResultTest.java
@@ -0,0 +1,66 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.beam.runners.spark.structuredstreaming;
+
+import static org.hamcrest.MatcherAssert.assertThat;
+import static org.hamcrest.Matchers.is;
+import static org.junit.Assert.assertFalse;
+
+import java.util.concurrent.CompletableFuture;
+import java.util.concurrent.Future;
+import java.util.concurrent.atomic.AtomicBoolean;
+import java.util.concurrent.atomic.AtomicInteger;
+import 
org.apache.beam.runners.spark.structuredstreaming.metrics.MetricsAccumulator;
+import org.apache.beam.sdk.PipelineResult.State;
+import org.junit.Test;
+import org.junit.runner.RunWith;
+import org.junit.runners.JUnit4;
+
+/** Tests for the cancel and wait semantics of {@link 
SparkStructuredStreamingPipelineResult}. */
+@RunWith(JUnit4.class)
+public class SparkStructuredStreamingPipelineResultTest {
+
+  private final AtomicInteger cancelSparkJobsCalls = new AtomicInteger();
+
+  private SparkStructuredStreamingPipelineResult result(Future<?> execution) {
+    return new SparkStructuredStreamingPipelineResult(
+        execution,
+        () -> null,
+        new MetricsAccumulator(),
+        new AtomicBoolean(),
+        cancelSparkJobsCalls::incrementAndGet);
+  }
+
+  @Test
+  public void testCancelRunsJobCancelHookOnce() throws Exception {
+    SparkStructuredStreamingPipelineResult result = result(new 
CompletableFuture<>());
+    assertThat(result.cancel(), is(State.CANCELLED));
+    assertThat(result.cancel(), is(State.CANCELLED));
+    assertThat(cancelSparkJobsCalls.get(), is(1));
+  }
+
+  @Test
+  public void testCancelIsAsynchronous() throws Exception {
+    CompletableFuture<Void> execution = new CompletableFuture<>();
+    SparkStructuredStreamingPipelineResult result = result(execution);
+    assertThat(result.cancel(), is(State.CANCELLED));
+    assertFalse(execution.isDone());
+    execution.completeExceptionally(new IllegalStateException("job 
cancelled"));
+    assertThat(result.waitUntilFinish(), is(State.CANCELLED));
+  }
+}
diff --git 
a/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/StructuredStreamingPipelineStateTest.java
 
b/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/StructuredStreamingPipelineStateTest.java
index b44df7bf101..647c4334ff1 100644
--- 
a/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/StructuredStreamingPipelineStateTest.java
+++ 
b/runners/spark/src/test/java/org/apache/beam/runners/spark/structuredstreaming/StructuredStreamingPipelineStateTest.java
@@ -20,10 +20,15 @@ package org.apache.beam.runners.spark.structuredstreaming;
 import static org.hamcrest.CoreMatchers.instanceOf;
 import static org.hamcrest.MatcherAssert.assertThat;
 import static org.hamcrest.Matchers.is;
+import static org.junit.Assert.assertFalse;
+import static org.junit.Assert.assertTrue;
 import static org.junit.Assert.fail;
 
 import java.io.Serializable;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.TimeUnit;
 import org.apache.beam.runners.spark.io.CreateStream;
+import 
org.apache.beam.runners.spark.structuredstreaming.translation.SparkSessionFactory;
 import org.apache.beam.sdk.Pipeline;
 import org.apache.beam.sdk.PipelineResult;
 import org.apache.beam.sdk.coders.StringUtf8Coder;
@@ -36,6 +41,8 @@ import org.apache.beam.sdk.transforms.ParDo;
 import org.apache.beam.sdk.transforms.SimpleFunction;
 import org.apache.beam.sdk.values.PBegin;
 import org.apache.beam.sdk.values.PCollection;
+import org.apache.spark.TaskContext;
+import org.apache.spark.sql.SparkSession;
 import org.joda.time.Duration;
 import org.junit.Ignore;
 import org.junit.Rule;
@@ -62,6 +69,33 @@ public class StructuredStreamingPipelineStateTest implements 
Serializable {
 
   private static final String FAILED_THE_BATCH_INTENTIONALLY = "Failed the 
batch intentionally";
 
+  private static final long DEADLINE_SECONDS = 60;
+
+  // Shared with the DoFn running in Spark's local executor threads, reset per 
test.
+  private static volatile CountDownLatch started = new CountDownLatch(1);
+
+  /** Signals started, then blocks until the task is killed. */
+  private static class BlockingDoFn extends DoFn<String, String> {
+    @ProcessElement
+    public void processElement(ProcessContext c) throws InterruptedException {
+      started.countDown();
+      while (!TaskContext.get().isInterrupted()) {
+        Thread.sleep(50);
+      }
+      c.output(c.element());
+    }
+  }
+
+  private SparkStructuredStreamingPipelineResult runBlockingPipeline() throws 
InterruptedException {
+    started = new CountDownLatch(1);
+    Pipeline pipeline = Pipeline.create(getBatchOptions());
+    pipeline.apply(Create.of("one", "two")).apply(ParDo.of(new 
BlockingDoFn()));
+    SparkStructuredStreamingPipelineResult result =
+        (SparkStructuredStreamingPipelineResult) pipeline.run();
+    assertTrue("DoFn did not start", started.await(DEADLINE_SECONDS, 
TimeUnit.SECONDS));
+    return result;
+  }
+
   private ParDo.SingleOutput<String, String> printParDo(final String prefix) {
     return ParDo.of(
         new DoFn<String, String>() {
@@ -151,6 +185,7 @@ public class StructuredStreamingPipelineStateTest 
implements Serializable {
     assertThat(result.getState(), is(PipelineResult.State.RUNNING));
 
     result.cancel();
+    assertThat(result.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
   }
 
   private void testCanceledPipeline(final 
SparkStructuredStreamingPipelineOptions options)
@@ -164,6 +199,7 @@ public class StructuredStreamingPipelineStateTest 
implements Serializable {
     result.cancel();
 
     assertThat(result.getState(), is(PipelineResult.State.CANCELLED));
+    assertThat(result.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
   }
 
   private void testRunningPipeline(final 
SparkStructuredStreamingPipelineOptions options)
@@ -177,6 +213,7 @@ public class StructuredStreamingPipelineStateTest 
implements Serializable {
     assertThat(result.getState(), is(PipelineResult.State.RUNNING));
 
     result.cancel();
+    assertThat(result.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
   }
 
   @Ignore("TODO: Reactivate with streaming.")
@@ -222,4 +259,39 @@ public class StructuredStreamingPipelineStateTest 
implements Serializable {
   public void testBatchPipelineTimeoutState() throws Exception {
     testTimeoutPipeline(getBatchOptions());
   }
+
+  @Test
+  public void testBatchCancelStopsRunningJob() throws Exception {
+    SparkStructuredStreamingPipelineResult result = runBlockingPipeline();
+    assertThat(result.cancel(), is(PipelineResult.State.CANCELLED));
+    assertThat(result.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
+    assertTrue("owned session not stopped", 
SparkSession.getDefaultSession().isEmpty());
+  }
+
+  @Test
+  public void testCancelKeepsSharedSession() throws Exception {
+    SparkSession session = 
SparkSessionFactory.sessionBuilder("local[1]").getOrCreate();
+    try {
+      SparkStructuredStreamingPipelineResult result = runBlockingPipeline();
+      assertThat(result.cancel(), is(PipelineResult.State.CANCELLED));
+      assertThat(result.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
+      assertFalse("shared session stopped", 
session.sparkContext().isStopped());
+    } finally {
+      session.stop();
+    }
+  }
+
+  /** The second pipeline shares the first session or creates a new one, both 
must end cleanly. */
+  @Test
+  public void testCancelFollowedImmediatelyBySecondPipeline() throws Exception 
{
+    SparkStructuredStreamingPipelineResult first = runBlockingPipeline();
+    assertThat(first.cancel(), is(PipelineResult.State.CANCELLED));
+    Pipeline secondPipeline = Pipeline.create(getBatchOptions());
+    secondPipeline.apply(Create.of("a", "b")).apply(printParDo("second"));
+    SparkStructuredStreamingPipelineResult second =
+        (SparkStructuredStreamingPipelineResult) secondPipeline.run();
+    assertThat(first.waitUntilFinish(), is(PipelineResult.State.CANCELLED));
+    assertThat(second.waitUntilFinish(), is(PipelineResult.State.DONE));
+    assertTrue("session not stopped", 
SparkSession.getDefaultSession().isEmpty());
+  }
 }

Reply via email to