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 c0f79b0251d [Spark][#36841] Add an opt-in Dataset-based backend for
portable pipelines (#40129)
c0f79b0251d is described below
commit c0f79b0251d178ac24dc819acabc5554536e322b
Author: Elia Liu <[email protected]>
AuthorDate: Wed Sep 23 05:36:27 2026 +1000
[Spark][#36841] Add an opt-in Dataset-based backend for portable pipelines
(#40129)
* [Spark][#36841] Add an opt-in Dataset-based backend for portable pipelines
Behind --useStructuredStreaming, SparkPipelineRunner translates the fused
Runner API pipeline into Spark Datasets. Impulse, Flatten, Reshuffle and
GroupByKey are Dataset operations, and executable stages run through the
existing Fn API bridge inside mapPartitions. Bounded pipelines run as batch
Datasets. Unbounded input, user state and timers are rejected at
translation.
The branch clears the streaming option so the metrics accumulator and
anything
else reading it agree with how the job runs.
The validatesPortableRunnerStructuredStreaming task runs the streaming
PortableValidatesRunner suite on this backend and is the exit gate for
making
it the default for portable streaming pipelines.
---
...Commit_Java_PVR_Spark4_StructuredStreaming.json | 3 +
.github/workflows/README.md | 1 +
...tCommit_Java_PVR_Spark4_StructuredStreaming.yml | 98 +++++
runners/spark/4/README.md | 6 +
runners/spark/job-server/spark_job_server.gradle | 10 +-
.../beam/runners/spark/SparkPipelineOptions.java | 9 +
.../beam/runners/spark/SparkPipelineRunner.java | 18 +-
.../runners/spark/metrics/MetricsAccumulator.java | 12 +-
.../SparkDatasetPortablePipelineTranslator.java | 379 ++++++++++++++++++
.../SparkDatasetTranslationContext.java | 126 ++++++
.../spark/SparkDatasetPortableExecutionTest.java | 178 +++++++++
...SparkDatasetPortablePipelineTranslatorTest.java | 442 +++++++++++++++++++++
.../SparkDatasetTranslationContextTest.java | 164 ++++++++
13 files changed, 1438 insertions(+), 8 deletions(-)
diff --git
a/.github/trigger_files/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.json
b/.github/trigger_files/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.json
new file mode 100644
index 00000000000..4c6e3309e30
--- /dev/null
+++
b/.github/trigger_files/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.json
@@ -0,0 +1,3 @@
+{
+ "comment": "Modify this file in a trivial way to cause this test suite to
run."
+}
diff --git a/.github/workflows/README.md b/.github/workflows/README.md
index 5d2d832f300..ff94d5d78b4 100644
--- a/.github/workflows/README.md
+++ b/.github/workflows/README.md
@@ -371,6 +371,7 @@ PostCommit Jobs run in a schedule against master branch and
generally do not get
| [ PostCommit Java PVR Spark Batch
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark_Batch.yml)
| N/A |`beam_PostCommit_Java_PVR_Spark_Batch.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark_Batch.yml?query=event%3Aschedule)
|
| [ PostCommit Java PVR Spark4 Batch
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Batch.yml)
| N/A |`beam_PostCommit_Java_PVR_Spark4_Batch.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Batch.yml?query=event%3Aschedule)
|
| [ PostCommit Java PVR Spark4 Streaming
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Streaming.yml)
| N/A |`beam_PostCommit_Java_PVR_Spark4_Streaming.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Streaming.yml?query=event
[...]
+| [ PostCommit Java PVR Spark4 StructuredStreaming
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml)
| N/A |`beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.json`|
[](https://github.com/apache/beam/actions/workflows/beam_Po
[...]
| [ PostCommit Java Tpcds Dataflow
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Dataflow.yml)
| N/A |`beam_PostCommit_Java_Tpcds_Dataflow.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Dataflow.yml?query=event%3Aschedule)
|
| [ PostCommit Java Tpcds Flink
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Flink.yml)
| N/A |`beam_PostCommit_Java_Tpcds_Flink.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Flink.yml?query=event%3Aschedule)
|
| [ PostCommit Java Tpcds Spark
](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Spark.yml)
| N/A |`beam_PostCommit_Java_Tpcds_Spark.json`|
[](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Spark.yml?query=event%3Aschedule)
|
diff --git
a/.github/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml
b/.github/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml
new file mode 100644
index 00000000000..33c8a3be79e
--- /dev/null
+++ b/.github/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml
@@ -0,0 +1,98 @@
+# 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.
+
+name: PostCommit Java PVR Spark4 StructuredStreaming
+
+on:
+ schedule:
+ - cron: '15 5/6 * * *'
+ pull_request_target:
+ paths: ['release/trigger_all_tests.json',
'.github/trigger_files/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.json']
+ workflow_dispatch:
+
+# This allows a subsequently queued workflow run to interrupt previous runs
+concurrency:
+ group: '${{ github.workflow }} @ ${{ github.event.pull_request.number ||
github.sha || github.head_ref || github.ref }}-${{ github.event.schedule ||
github.event.comment.id || github.event.sender.login }}'
+ cancel-in-progress: true
+
+#Setting explicit permissions for the action to avoid the default permissions
which are `write-all` in case of pull_request_target event
+permissions:
+ actions: write
+ pull-requests: write
+ checks: write
+ contents: read
+ deployments: read
+ id-token: none
+ issues: write
+ discussions: read
+ packages: read
+ pages: read
+ repository-projects: read
+ security-events: read
+ statuses: read
+
+env:
+ DEVELOCITY_ACCESS_KEY: ${{ secrets.DEVELOCITY_ACCESS_KEY }}
+ GRADLE_ENTERPRISE_CACHE_USERNAME: ${{ secrets.GE_CACHE_USERNAME }}
+ GRADLE_ENTERPRISE_CACHE_PASSWORD: ${{ secrets.GE_CACHE_PASSWORD }}
+
+jobs:
+ beam_PostCommit_Java_PVR_Spark4_StructuredStreaming:
+ name: ${{ matrix.job_name }} (${{ matrix.job_phrase }})
+ runs-on: [self-hosted, ubuntu-24.04, main]
+ timeout-minutes: 180
+ strategy:
+ matrix:
+ job_name: [beam_PostCommit_Java_PVR_Spark4_StructuredStreaming]
+ job_phrase: [Run Java Spark v4 PortableValidatesRunner
StructuredStreaming]
+ if: |
+ github.event_name == 'workflow_dispatch' ||
+ github.event_name == 'pull_request_target' ||
+ (github.event_name == 'schedule' && github.repository == 'apache/beam')
||
+ github.event.comment.body == 'Run Java Spark v4 PortableValidatesRunner
StructuredStreaming'
+ steps:
+ - uses: actions/checkout@v7
+ with:
+ persist-credentials: false
+ - name: Setup repository
+ uses: ./.github/actions/setup-action
+ with:
+ comment_phrase: ${{ matrix.job_phrase }}
+ github_token: ${{ secrets.GITHUB_TOKEN }}
+ github_job: ${{ matrix.job_name }} (${{ matrix.job_phrase }})
+ - name: Setup environment
+ uses: ./.github/actions/setup-environment-action
+ with:
+ java-version: 17
+ - name: run PostCommit Java PortableValidatesRunner Spark4
StructuredStreaming script
+ uses: ./.github/actions/gradle-command-self-hosted-action
+ with:
+ gradle-command:
:runners:spark:4:job-server:validatesPortableRunnerStructuredStreaming
+ - name: Archive JUnit Test Results
+ uses: actions/upload-artifact@v7
+ if: ${{ !success() }}
+ with:
+ name: JUnit Test Results
+ path: "**/build/reports/tests/"
+ - name: Publish JUnit Test Results
+ uses: EnricoMi/publish-unit-test-result-action@v2
+ if: always()
+ with:
+ commit: '${{ env.prsha || env.GITHUB_SHA }}'
+ comment_mode: ${{ github.event_name == 'issue_comment' && 'always'
|| 'off' }}
+ files: '**/build/test-results/**/*.xml'
+ large_files: true
diff --git a/runners/spark/4/README.md b/runners/spark/4/README.md
index 371a4451265..02ba1f0b5e6 100644
--- a/runners/spark/4/README.md
+++ b/runners/spark/4/README.md
@@ -36,6 +36,12 @@ runner code.
Batch only. Streaming is tracked in
[#36841](https://github.com/apache/beam/issues/36841).
+The portable job server can run bounded pipelines on the Dataset-based backend
+with `--useStructuredStreaming`. That path is experimental and does not support
+unbounded input, state or timers yet. The
+`validatesPortableRunnerStructuredStreaming` task runs the streaming
+PortableValidatesRunner suite against it.
+
## Known issues
### `StackOverflowError` from `slf4j-jdk14` on the runtime classpath
diff --git a/runners/spark/job-server/spark_job_server.gradle
b/runners/spark/job-server/spark_job_server.gradle
index 0166f065401..637ac98375d 100644
--- a/runners/spark/job-server/spark_job_server.gradle
+++ b/runners/spark/job-server/spark_job_server.gradle
@@ -94,8 +94,11 @@ def sickbayTests = [
'org.apache.beam.sdk.transforms.ReshuffleTest.testReshufflePreservesMetadata',
]
-def portableValidatesRunnerTask(String name, boolean streaming, boolean
docker, ArrayList<String> sickbayTests) {
+def portableValidatesRunnerTask(String name, boolean streaming, boolean
docker, ArrayList<String> sickbayTests, boolean structuredStreaming = false) {
def pipelineOptions = []
+ if (structuredStreaming) {
+ pipelineOptions += "--useStructuredStreaming"
+ }
def testCategories
def testFilter
@@ -243,6 +246,9 @@ def portableValidatesRunnerTask(String name, boolean
streaming, boolean docker,
project.ext.validatesPortableRunnerDocker=
portableValidatesRunnerTask("Docker", false, true, sickbayTests)
project.ext.validatesPortableRunnerBatch =
portableValidatesRunnerTask("Batch", false, false, sickbayTests)
project.ext.validatesPortableRunnerStreaming =
portableValidatesRunnerTask("Streaming", true, false, sickbayTests)
+// Structured Streaming variant of the streaming suite: same tests and
exclusions, run on the
+// Dataset-based backend. It is the exit gate for making that backend the
streaming default.
+project.ext.validatesPortableRunnerStructuredStreaming =
portableValidatesRunnerTask("StructuredStreaming", true, false, sickbayTests,
true)
tasks.register("validatesPortableRunner") {
dependsOn validatesPortableRunnerDocker
@@ -284,7 +290,7 @@ def sparkJobServerJvmArgs() {
}
// TestPortableRunner starts SparkJobServerDriver in-process in the test JVM.
-['validatesPortableRunnerDocker', 'validatesPortableRunnerBatch',
'validatesPortableRunnerStreaming'].each { taskName ->
+['validatesPortableRunnerDocker', 'validatesPortableRunnerBatch',
'validatesPortableRunnerStreaming',
'validatesPortableRunnerStructuredStreaming'].each { taskName ->
tasks.named(taskName) {
jvmArgs += sparkJobServerJvmArgs()
}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineOptions.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineOptions.java
index 2ad149431d2..f2d5b2c1272 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineOptions.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineOptions.java
@@ -86,4 +86,13 @@ public interface SparkPipelineOptions extends
SparkCommonPipelineOptions {
boolean isCacheDisabled();
void setCacheDisabled(boolean value);
+
+ @Description(
+ "Run portable pipelines on the Dataset-based backend. Experimental.
Unbounded input, user"
+ + " state and timers are not supported on this backend yet, see"
+ + " https://github.com/apache/beam/issues/36841.")
+ @Default.Boolean(false)
+ boolean getUseStructuredStreaming();
+
+ void setUseStructuredStreaming(boolean value);
}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineRunner.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineRunner.java
index 91a94896b89..51fe26d889f 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineRunner.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/SparkPipelineRunner.java
@@ -35,6 +35,7 @@ import
org.apache.beam.runners.jobsubmission.PortablePipelineRunner;
import org.apache.beam.runners.spark.metrics.MetricsAccumulator;
import
org.apache.beam.runners.spark.translation.SparkBatchPortablePipelineTranslator;
import org.apache.beam.runners.spark.translation.SparkContextFactory;
+import
org.apache.beam.runners.spark.translation.SparkDatasetPortablePipelineTranslator;
import
org.apache.beam.runners.spark.translation.SparkPortablePipelineTranslator;
import
org.apache.beam.runners.spark.translation.SparkStreamingPortablePipelineTranslator;
import
org.apache.beam.runners.spark.translation.SparkStreamingTranslationContext;
@@ -81,8 +82,14 @@ public class SparkPipelineRunner implements
PortablePipelineRunner {
@Override
public PortablePipelineResult run(RunnerApi.Pipeline pipeline, JobInfo
jobInfo) {
SparkPortablePipelineTranslator translator;
- boolean isStreaming = pipelineOptions.isStreaming() ||
hasUnboundedPCollections(pipeline);
- if (isStreaming) {
+ boolean useStructuredStreaming =
pipelineOptions.getUseStructuredStreaming();
+ // The Dataset backend never uses the DStream translator or a streaming
context.
+ boolean useDStreams =
+ !useStructuredStreaming
+ && (pipelineOptions.isStreaming() ||
hasUnboundedPCollections(pipeline));
+ if (useStructuredStreaming) {
+ translator = new SparkDatasetPortablePipelineTranslator();
+ } else if (useDStreams) {
translator = new SparkStreamingPortablePipelineTranslator();
} else {
translator = new SparkBatchPortablePipelineTranslator();
@@ -112,9 +119,10 @@ public class SparkPipelineRunner implements
PortablePipelineRunner {
PortablePipelineResult result;
final JavaSparkContext jsc =
SparkContextFactory.getSparkContext(pipelineOptions);
- // Initialize accumulators.
+ // Initialize accumulators. Only the DStream streaming path uses the
metrics checkpoint.
MetricsEnvironment.setMetricsSupported(true);
- MetricsAccumulator.init(pipelineOptions, jsc);
+ MetricsAccumulator.init(
+ pipelineOptions, jsc, !useStructuredStreaming &&
pipelineOptions.isStreaming());
final SparkTranslationContext context =
translator.createTranslationContext(jsc, pipelineOptions, jobInfo);
@@ -127,7 +135,7 @@ public class SparkPipelineRunner implements
PortablePipelineRunner {
LOG.info("Running job {} on Spark master {}", jobInfo.jobId(),
jsc.master());
- if (isStreaming) {
+ if (useDStreams) {
final JavaStreamingContext jssc =
((SparkStreamingTranslationContext) context).getStreamingContext();
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/metrics/MetricsAccumulator.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/metrics/MetricsAccumulator.java
index 612d71b1aea..0e06b5f66e3 100644
---
a/runners/spark/src/main/java/org/apache/beam/runners/spark/metrics/MetricsAccumulator.java
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/metrics/MetricsAccumulator.java
@@ -55,11 +55,21 @@ public class MetricsAccumulator {
/** Init metrics accumulator if it has not been initiated. This method is
idempotent. */
public static void init(SparkPipelineOptions opts, JavaSparkContext jsc) {
+ init(opts, jsc, opts.isStreaming());
+ }
+
+ /**
+ * Init metrics accumulator if it has not been initiated. This method is
idempotent. With {@code
+ * useCheckpoint} set, the value is recovered from the metrics checkpoint
under the checkpoint
+ * directory of {@code opts}, and {@link
AccumulatorCheckpointingSparkListener} writes it back
+ * there. The DStream streaming path is the one using that checkpoint.
+ */
+ public static void init(SparkPipelineOptions opts, JavaSparkContext jsc,
boolean useCheckpoint) {
if (instance == null) {
synchronized (MetricsAccumulator.class) {
if (instance == null) {
Optional<CheckpointDir> maybeCheckpointDir =
- opts.isStreaming()
+ useCheckpoint
? Optional.of(new CheckpointDir(opts.getCheckpointDir()))
: Optional.absent();
MetricsContainerStepMap metricsContainerStepMap = new
SparkMetricsContainerStepMap();
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslator.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslator.java
new file mode 100644
index 00000000000..b58920ab198
--- /dev/null
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslator.java
@@ -0,0 +1,379 @@
+/*
+ * 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.translation;
+
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.createOutputMap;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getInputId;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getOutputId;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getWindowedValueCoder;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getWindowingStrategy;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.hasUnboundedPCollections;
+import static org.apache.spark.sql.functions.col;
+
+import java.io.IOException;
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.HashMap;
+import java.util.Iterator;
+import java.util.List;
+import java.util.Map;
+import java.util.Set;
+import org.apache.beam.model.pipeline.v1.RunnerApi;
+import
org.apache.beam.model.pipeline.v1.RunnerApi.ExecutableStagePayload.SideInputId;
+import org.apache.beam.runners.core.SystemReduceFn;
+import org.apache.beam.runners.fnexecution.provisioning.JobInfo;
+import org.apache.beam.runners.spark.SparkPipelineOptions;
+import org.apache.beam.runners.spark.coders.CoderHelpers;
+import org.apache.beam.runners.spark.metrics.MetricsAccumulator;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.helpers.EncoderHelpers;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.transforms.join.RawUnionValue;
+import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
+import org.apache.beam.sdk.util.construction.PTransformTranslation;
+import org.apache.beam.sdk.util.construction.graph.ExecutableStage;
+import org.apache.beam.sdk.util.construction.graph.PipelineNode.PTransformNode;
+import org.apache.beam.sdk.util.construction.graph.QueryablePipeline;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowedValues;
+import org.apache.beam.sdk.values.WindowedValues.WindowedValueCoder;
+import org.apache.beam.sdk.values.WindowingStrategy;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.annotations.VisibleForTesting;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.BiMap;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Iterators;
+import org.apache.spark.api.java.JavaSparkContext;
+import org.apache.spark.api.java.function.FlatMapGroupsFunction;
+import org.apache.spark.api.java.function.MapFunction;
+import org.apache.spark.api.java.function.MapPartitionsFunction;
+import org.apache.spark.broadcast.Broadcast;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Encoder;
+import org.apache.spark.sql.Encoders;
+import org.apache.spark.sql.TypedColumn;
+import scala.Tuple2;
+
+/**
+ * Translates a portable pipeline into Spark Dataset operations.
+ *
+ * <p>Executable stages run through the Fn API bridge of {@link
SparkExecutableStageFunction} inside
+ * {@code mapPartitions}. Side inputs are collected and broadcast. Unbounded
input, user state and
+ * timers are not supported yet and fail at translation.
+ */
+@SuppressWarnings({
+ "rawtypes", // TODO(https://github.com/apache/beam/issues/20447)
+ "unchecked",
+ "nullness" // TODO(https://github.com/apache/beam/issues/20497)
+})
+public class SparkDatasetPortablePipelineTranslator
+ implements SparkPortablePipelineTranslator<SparkDatasetTranslationContext>
{
+
+ private final ImmutableMap<String, PTransformTranslator>
urnToTransformTranslator;
+
+ interface PTransformTranslator {
+ void translate(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context);
+ }
+
+ public SparkDatasetPortablePipelineTranslator() {
+ ImmutableMap.Builder<String, PTransformTranslator> translatorMap =
ImmutableMap.builder();
+ translatorMap.put(
+ PTransformTranslation.IMPULSE_TRANSFORM_URN,
+ SparkDatasetPortablePipelineTranslator::translateImpulse);
+ translatorMap.put(
+ PTransformTranslation.GROUP_BY_KEY_TRANSFORM_URN,
+ SparkDatasetPortablePipelineTranslator::translateGroupByKey);
+ translatorMap.put(
+ ExecutableStage.URN,
SparkDatasetPortablePipelineTranslator::translateExecutableStage);
+ translatorMap.put(
+ PTransformTranslation.FLATTEN_TRANSFORM_URN,
+ SparkDatasetPortablePipelineTranslator::translateFlatten);
+ translatorMap.put(
+ PTransformTranslation.RESHUFFLE_URN,
+ SparkDatasetPortablePipelineTranslator::translateReshuffle);
+ this.urnToTransformTranslator = translatorMap.build();
+ }
+
+ @Override
+ public Set<String> knownUrns() {
+ return urnToTransformTranslator.keySet();
+ }
+
+ @Override
+ public void translate(RunnerApi.Pipeline pipeline,
SparkDatasetTranslationContext context) {
+ if (hasUnboundedPCollections(pipeline)) {
+ throw new UnsupportedOperationException(
+ "The Dataset-based portable Spark runner does not support unbounded
input yet, see"
+ + " https://github.com/apache/beam/issues/36841.");
+ }
+ QueryablePipeline p =
+ QueryablePipeline.forTransforms(
+ pipeline.getRootTransformIdsList(), pipeline.getComponents());
+ for (PTransformNode transformNode : p.getTopologicallyOrderedTransforms())
{
+ for (String inputId :
transformNode.getTransform().getInputsMap().values()) {
+ context.addConsumer(inputId);
+ }
+ }
+ for (PTransformNode transformNode : p.getTopologicallyOrderedTransforms())
{
+ urnToTransformTranslator
+ .getOrDefault(
+ transformNode.getTransform().getSpec().getUrn(),
+ SparkDatasetPortablePipelineTranslator::urnNotFound)
+ .translate(transformNode, pipeline, context);
+ }
+ }
+
+ @Override
+ public SparkDatasetTranslationContext createTranslationContext(
+ JavaSparkContext jsc, SparkPipelineOptions options, JobInfo jobInfo) {
+ return new SparkDatasetTranslationContext(jsc, options, jobInfo);
+ }
+
+ private static void urnNotFound(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ throw new IllegalArgumentException(
+ String.format(
+ "Transform %s has unknown URN %s",
+ transformNode.getId(),
transformNode.getTransform().getSpec().getUrn()));
+ }
+
+ @VisibleForTesting
+ static void translateImpulse(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ String outputId = getOutputId(transformNode);
+ Dataset<WindowedValue<byte[]>> dataset =
+ context
+ .getSparkSession()
+ .createDataset(
+
Collections.singletonList(WindowedValues.valueInGlobalWindow(new byte[0])),
+ context.windowedEncoder(outputId, pipeline.getComponents()));
+ context.putDataset(outputId, dataset);
+ }
+
+ @VisibleForTesting
+ static <InputT, SideInputT> void translateExecutableStage(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ RunnerApi.ExecutableStagePayload stagePayload;
+ try {
+ stagePayload =
+ RunnerApi.ExecutableStagePayload.parseFrom(
+ transformNode.getTransform().getSpec().getPayload());
+ } catch (IOException e) {
+ throw new RuntimeException(e);
+ }
+ if (stagePayload.getUserStatesCount() > 0 || stagePayload.getTimersCount()
> 0) {
+ throw new UnsupportedOperationException(
+ String.format(
+ "Stage %s uses state or timers, which the Dataset-based portable
Spark runner does"
+ + " not support yet, see
https://github.com/apache/beam/issues/20396 and"
+ + " https://github.com/apache/beam/issues/20397.",
+ transformNode.getId()));
+ }
+ RunnerApi.Components components = pipeline.getComponents();
+ String inputId = stagePayload.getInput();
+ Dataset<WindowedValue<InputT>> input = context.getDataset(inputId);
+ Map<String, String> outputs = transformNode.getTransform().getOutputsMap();
+ BiMap<String, Integer> outputMap = createOutputMap(outputs.values());
+ Coder windowCoder = getWindowingStrategy(inputId,
components).getWindowFn().windowCoder();
+
+ SparkExecutableStageFunction<InputT, SideInputT> stageFunction =
+ new SparkExecutableStageFunction<>(
+ context.getSerializableOptions(),
+ stagePayload,
+ context.jobInfo,
+ outputMap,
+ SparkExecutableStageContextFactory.getInstance(),
+ broadcastSideInputs(stagePayload, context),
+ MetricsAccumulator.getInstance(),
+ windowCoder,
+ getWindowedValueCoder(inputId, components),
+ true);
+
+ if (outputs.isEmpty()) {
+ // Fusion can leave a stage without runner-visible output. It still has
to run, so it
+ // becomes a leaf that emits nothing.
+ Dataset<WindowedValue<InputT>> sink =
+ input.mapPartitions(
+ (MapPartitionsFunction<WindowedValue<InputT>,
WindowedValue<InputT>>)
+ elements -> {
+ Iterator<RawUnionValue> results =
stageFunction.call(elements);
+ while (results.hasNext()) {
+ results.next();
+ }
+ return Collections.emptyIterator();
+ },
+ input.encoder());
+ context.putDataset(String.format("EmptyOutputSink_%d",
context.nextSinkId()), sink);
+ return;
+ }
+
+ // One encoder per output, in union tag order.
+ List<Encoder<WindowedValue<Object>>> encoders =
+ new ArrayList<>(Collections.nCopies(outputMap.size(), null));
+ for (Map.Entry<String, Integer> output : outputMap.entrySet()) {
+ encoders.set(output.getValue(), context.windowedEncoder(output.getKey(),
components));
+ }
+ Dataset<Tuple2<Integer, WindowedValue<Object>>> staged =
+ input.mapPartitions(
+ (MapPartitionsFunction<WindowedValue<InputT>, Tuple2<Integer,
WindowedValue<Object>>>)
+ elements -> tagged(stageFunction.call(elements)),
+ EncoderHelpers.oneOfEncoder(encoders));
+ boolean staging = outputs.size() > 1;
+ if (staging) {
+ // Every output is a projection of the same stage run. Persist so the
stage runs once.
+ staged = staged.persist(context.getStorageLevel());
+ }
+ for (Map.Entry<String, Integer> output : outputMap.entrySet()) {
+ int tag = output.getValue();
+ TypedColumn<Tuple2<Integer, WindowedValue<Object>>,
WindowedValue<Object>> column =
+ (TypedColumn) col(Integer.toString(tag)).as(encoders.get(tag));
+ // A projection of persisted rows is not cached again.
+ context.putDataset(
+ output.getKey(), staged.filter(column.isNotNull()).select(column),
!staging);
+ }
+ }
+
+ private static Iterator<Tuple2<Integer, WindowedValue<Object>>> tagged(
+ Iterator<RawUnionValue> values) {
+ return Iterators.transform(
+ values,
+ value -> new Tuple2<>(value.getUnionTag(), (WindowedValue<Object>)
value.getValue()));
+ }
+
+ /** Collects each side input of a stage and broadcasts its encoded elements.
*/
+ private static <SideInputT>
+ ImmutableMap<String, Tuple2<Broadcast<List<byte[]>>,
WindowedValueCoder<SideInputT>>>
+ broadcastSideInputs(
+ RunnerApi.ExecutableStagePayload stagePayload,
+ SparkDatasetTranslationContext context) {
+ Map<String, Tuple2<Broadcast<List<byte[]>>,
WindowedValueCoder<SideInputT>>> broadcasts =
+ new HashMap<>();
+ RunnerApi.Components components = stagePayload.getComponents();
+ for (SideInputId sideInputId : stagePayload.getSideInputsList()) {
+ String collectionId =
+ components
+ .getTransformsOrThrow(sideInputId.getTransformId())
+ .getInputsOrThrow(sideInputId.getLocalName());
+ if (broadcasts.containsKey(collectionId)) {
+ continue;
+ }
+ WindowedValueCoder<SideInputT> coder =
getWindowedValueCoder(collectionId, components);
+ Dataset<WindowedValue<SideInputT>> dataset =
context.getDataset(collectionId);
+ List<byte[]> bytes =
+ new ArrayList<>(
+ dataset
+ .map(
+ (MapFunction<WindowedValue<SideInputT>, byte[]>)
+ value -> CoderHelpers.toByteArray(value, coder),
+ Encoders.BINARY())
+ .collectAsList());
+ broadcasts.put(collectionId, new
Tuple2<>(context.getSparkContext().broadcast(bytes), coder));
+ }
+ return ImmutableMap.copyOf(broadcasts);
+ }
+
+ @VisibleForTesting
+ static <K, V> void translateGroupByKey(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ RunnerApi.Components components = pipeline.getComponents();
+ String inputId = getInputId(transformNode);
+ String outputId = getOutputId(transformNode);
+ Dataset<WindowedValue<KV<K, V>>> input = context.getDataset(inputId);
+ WindowedValueCoder<KV<K, V>> inputCoder = getWindowedValueCoder(inputId,
components);
+ KvCoder<K, V> kvCoder = (KvCoder<K, V>) inputCoder.getValueCoder();
+ Coder<K> keyCoder = kvCoder.getKeyCoder();
+ WindowingStrategy<?, BoundedWindow> windowingStrategy =
+ getWindowingStrategy(inputId, components);
+
+ // Batch semantics: all values of a key are present, so every window of
the key can close.
+ SparkGroupAlsoByWindowViaOutputBufferFn<K, V, BoundedWindow>
groupAlsoByWindow =
+ new SparkGroupAlsoByWindowViaOutputBufferFn<>(
+ windowingStrategy,
+ new TranslationUtils.InMemoryStateInternalsFactory<>(),
+ SystemReduceFn.buffering(kvCoder.getValueCoder()),
+ context.getSerializableOptions());
+
+ Dataset<WindowedValue<KV<K, Iterable<V>>>> grouped =
+ input
+ .groupByKey(
+ (MapFunction<WindowedValue<KV<K, V>>, byte[]>)
+ value ->
CoderHelpers.toByteArray(value.getValue().getKey(), keyCoder),
+ Encoders.BINARY())
+ .flatMapGroups(
+ (FlatMapGroupsFunction<
+ byte[], WindowedValue<KV<K, V>>, WindowedValue<KV<K,
Iterable<V>>>>)
+ (keyBytes, values) -> {
+ K key = CoderHelpers.fromByteArray(keyBytes, keyCoder);
+ List<WindowedValue<V>> windowedValues = new
ArrayList<>();
+ while (values.hasNext()) {
+ WindowedValue<KV<K, V>> value = values.next();
+
windowedValues.add(value.withValue(value.getValue().getValue()));
+ }
+ return groupAlsoByWindow.call(
+ KV.<K, Iterable<WindowedValue<V>>>of(key,
windowedValues));
+ },
+ context.windowedEncoder(outputId, components));
+ context.putDataset(outputId, grouped);
+ }
+
+ @VisibleForTesting
+ static <T> void translateFlatten(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ RunnerApi.Components components = pipeline.getComponents();
+ String outputId = getOutputId(transformNode);
+ WindowedValueCoder<T> outputCoder = getWindowedValueCoder(outputId,
components);
+ Encoder<WindowedValue<T>> outputEncoder =
context.windowedEncoder(outputId, components);
+ Dataset<WindowedValue<T>> result = null;
+ for (String inputId :
transformNode.getTransform().getInputsMap().values()) {
+ Dataset<WindowedValue<T>> input = context.getDataset(inputId);
+ if (!getWindowedValueCoder(inputId, components).equals(outputCoder)) {
+ // Re-encode so every branch of the union shares the output schema.
+ input = input.map((MapFunction<WindowedValue<T>, WindowedValue<T>>) v
-> v, outputEncoder);
+ }
+ result = result == null ? input : result.union(input);
+ }
+ if (result == null) {
+ result = context.getSparkSession().emptyDataset(outputEncoder);
+ }
+ context.putDataset(outputId, result);
+ }
+
+ @VisibleForTesting
+ static <T> void translateReshuffle(
+ PTransformNode transformNode,
+ RunnerApi.Pipeline pipeline,
+ SparkDatasetTranslationContext context) {
+ Dataset<WindowedValue<T>> input =
context.getDataset(getInputId(transformNode));
+ context.putDataset(
+ getOutputId(transformNode),
+ input.repartition(context.getSparkContext().defaultParallelism()));
+ }
+}
diff --git
a/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContext.java
b/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContext.java
new file mode 100644
index 00000000000..d58fd62cd7e
--- /dev/null
+++
b/runners/spark/src/main/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContext.java
@@ -0,0 +1,126 @@
+/*
+ * 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.translation;
+
+import java.util.HashMap;
+import java.util.LinkedHashSet;
+import java.util.Map;
+import java.util.Set;
+import org.apache.beam.model.pipeline.v1.RunnerApi;
+import org.apache.beam.runners.fnexecution.provisioning.JobInfo;
+import org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils;
+import org.apache.beam.runners.spark.SparkPipelineOptions;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.EvaluationContext;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.helpers.EncoderHelpers;
+import org.apache.beam.sdk.coders.Coder;
+import org.apache.beam.sdk.transforms.windowing.BoundedWindow;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.spark.api.java.JavaSparkContext;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Encoder;
+import org.apache.spark.sql.SparkSession;
+import org.apache.spark.storage.StorageLevel;
+
+/**
+ * Translation context of the Dataset-based portable backend. Keeps one {@link
Dataset} per
+ * translated PCollection and evaluates the ones no transform consumed in
{@link #computeOutputs()}.
+ */
+@SuppressWarnings({
+ "rawtypes", // TODO(https://github.com/apache/beam/issues/20447)
+ "unchecked",
+ "nullness" // TODO(https://github.com/apache/beam/issues/20497)
+})
+public class SparkDatasetTranslationContext extends SparkTranslationContext {
+ private final SparkSession session;
+ private final StorageLevel storageLevel;
+ private final boolean cacheDisabled;
+ private final Map<String, Integer> consumers = new HashMap<>();
+ private final Map<String, Dataset> datasets = new HashMap<>();
+ private final Set<String> leaves = new LinkedHashSet<>();
+ private final Map<Coder<?>, Encoder<?>> encoders = new HashMap<>();
+
+ public SparkDatasetTranslationContext(
+ JavaSparkContext jsc, SparkPipelineOptions options, JobInfo jobInfo) {
+ super(jsc, options, jobInfo);
+ // The builder attaches to the SparkContext that SparkContextFactory
already created.
+ this.session = SparkSession.builder().getOrCreate();
+ this.storageLevel = StorageLevel.fromString(options.getStorageLevel());
+ this.cacheDisabled = options.isCacheDisabled();
+ }
+
+ public SparkSession getSparkSession() {
+ return session;
+ }
+
+ public StorageLevel getStorageLevel() {
+ return storageLevel;
+ }
+
+ /** Records one more transform reading {@code pCollectionId}. */
+ void addConsumer(String pCollectionId) {
+ consumers.merge(pCollectionId, 1, Integer::sum);
+ }
+
+ /** Registers the Dataset of a PCollection. Datasets read by several
transforms are persisted. */
+ public <T> void putDataset(String pCollectionId, Dataset<WindowedValue<T>>
dataset) {
+ putDataset(pCollectionId, dataset, true);
+ }
+
+ /**
+ * Registers the Dataset of a PCollection. Pass {@code cache} as false for a
Dataset that is a
+ * projection of an already persisted one, so the same rows are not cached
twice.
+ */
+ public <T> void putDataset(
+ String pCollectionId, Dataset<WindowedValue<T>> dataset, boolean cache) {
+ if (cache && !cacheDisabled && consumers.getOrDefault(pCollectionId, 0) >
1) {
+ dataset = dataset.persist(storageLevel);
+ }
+ datasets.put(pCollectionId, dataset);
+ leaves.add(pCollectionId);
+ }
+
+ /** Returns the Dataset of a PCollection and marks it as consumed. */
+ public <T> Dataset<WindowedValue<T>> getDataset(String pCollectionId) {
+ leaves.remove(pCollectionId);
+ return datasets.get(pCollectionId);
+ }
+
+ /** Encoder of the windowed values of a PCollection, derived from its wire
coder. */
+ public <T> Encoder<WindowedValue<T>> windowedEncoder(
+ String pCollectionId, RunnerApi.Components components) {
+ Coder<T> valueCoder =
+ PipelineTranslatorUtils.<T>getWindowedValueCoder(pCollectionId,
components).getValueCoder();
+ Coder<? extends BoundedWindow> windowCoder =
+ PipelineTranslatorUtils.getWindowingStrategy(pCollectionId, components)
+ .getWindowFn()
+ .windowCoder();
+ return EncoderHelpers.windowedValueEncoder(encoderOf(valueCoder),
encoderOf(windowCoder));
+ }
+
+ private <T> Encoder<T> encoderOf(Coder<T> coder) {
+ return (Encoder<T>) encoders.computeIfAbsent(coder, c ->
EncoderHelpers.encoderFor(coder));
+ }
+
+ /** Evaluates every Dataset no transform consumed. */
+ @Override
+ public void computeOutputs() {
+ for (String leaf : leaves) {
+ EvaluationContext.evaluate(leaf, datasets.get(leaf));
+ }
+ }
+}
diff --git
a/runners/spark/src/test/java/org/apache/beam/runners/spark/SparkDatasetPortableExecutionTest.java
b/runners/spark/src/test/java/org/apache/beam/runners/spark/SparkDatasetPortableExecutionTest.java
new file mode 100644
index 00000000000..90375d7f6e2
--- /dev/null
+++
b/runners/spark/src/test/java/org/apache/beam/runners/spark/SparkDatasetPortableExecutionTest.java
@@ -0,0 +1,178 @@
+/*
+ * 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;
+
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertTrue;
+
+import java.io.Serializable;
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.List;
+import java.util.concurrent.CopyOnWriteArrayList;
+import java.util.concurrent.Executors;
+import java.util.concurrent.TimeUnit;
+import org.apache.beam.model.jobmanagement.v1.JobApi.JobState;
+import org.apache.beam.model.pipeline.v1.RunnerApi;
+import org.apache.beam.runners.jobsubmission.JobInvocation;
+import org.apache.beam.sdk.Pipeline;
+import org.apache.beam.sdk.coders.BigEndianLongCoder;
+import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.coders.StringUtf8Coder;
+import org.apache.beam.sdk.options.PipelineOptions;
+import org.apache.beam.sdk.options.PipelineOptionsFactory;
+import org.apache.beam.sdk.options.PortablePipelineOptions;
+import org.apache.beam.sdk.testing.CrashingRunner;
+import org.apache.beam.sdk.testing.PAssert;
+import org.apache.beam.sdk.transforms.DoFn;
+import org.apache.beam.sdk.transforms.Flatten;
+import org.apache.beam.sdk.transforms.GroupByKey;
+import org.apache.beam.sdk.transforms.Impulse;
+import org.apache.beam.sdk.transforms.ParDo;
+import org.apache.beam.sdk.transforms.WithKeys;
+import org.apache.beam.sdk.util.construction.Environments;
+import org.apache.beam.sdk.util.construction.PipelineTranslation;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.PCollection;
+import org.apache.beam.sdk.values.PCollectionList;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.ListeningExecutorService;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.util.concurrent.MoreExecutors;
+import org.junit.AfterClass;
+import org.junit.BeforeClass;
+import org.junit.Test;
+import org.junit.runner.RunWith;
+import org.junit.runners.JUnit4;
+
+/**
+ * Runs one portable pipeline end to end on the Dataset-based backend through
the job invoker, with
+ * {@code --streaming} and {@code --useStructuredStreaming} as the {@code
+ * validatesPortableRunnerStructuredStreaming} task sets them. The translator
and its context have
+ * their own unit tests.
+ */
+@RunWith(JUnit4.class)
+public class SparkDatasetPortableExecutionTest implements Serializable {
+ private static ListeningExecutorService executor;
+
+ @BeforeClass
+ public static void setUp() {
+ executor =
MoreExecutors.listeningDecorator(Executors.newFixedThreadPool(1));
+ }
+
+ @AfterClass
+ public static void tearDown() throws InterruptedException {
+ executor.shutdown();
+ executor.awaitTermination(10, TimeUnit.SECONDS);
+ executor = null;
+ }
+
+ @Test(timeout = 180_000)
+ public void boundedPipelineRunsOnDatasets() throws Exception {
+ SparkPipelineOptions options = options();
+ Pipeline p = Pipeline.create(options);
+ PCollection<String> words =
+ p.apply("impulse", Impulse.create())
+ .apply(
+ "create",
+ ParDo.of(
+ new DoFn<byte[], String>() {
+ @ProcessElement
+ public void process(ProcessContext ctxt) {
+ ctxt.output("zero");
+ ctxt.output("one");
+ ctxt.output("two");
+ }
+ }));
+ PCollection<String> more =
+ p.apply("impulse2", Impulse.create())
+ .apply(
+ "create2",
+ ParDo.of(
+ new DoFn<byte[], String>() {
+ @ProcessElement
+ public void process(ProcessContext ctxt) {
+ ctxt.output("three");
+ }
+ }));
+ PCollection<String> result =
+ PCollectionList.of(words)
+ .and(more)
+ .apply("flatten", Flatten.pCollections())
+ .apply(
+ "len",
+ ParDo.of(
+ new DoFn<String, Long>() {
+ @ProcessElement
+ public void process(ProcessContext ctxt) {
+ ctxt.output((long) ctxt.element().length());
+ }
+ }))
+ .apply("addKeys", WithKeys.of("foo"))
+ // Use some unknown coders
+ .setCoder(KvCoder.of(StringUtf8Coder.of(),
BigEndianLongCoder.of()))
+ .apply("gbk", GroupByKey.create())
+ .apply(
+ "format",
+ ParDo.of(
+ new DoFn<KV<String, Iterable<Long>>, String>() {
+ @ProcessElement
+ public void process(ProcessContext ctxt) {
+ // The order of grouped values is not defined, so sort
before comparing.
+ List<Long> values = new ArrayList<>();
+ ctxt.element().getValue().forEach(values::add);
+ Collections.sort(values);
+ ctxt.output(ctxt.element().getKey() + ":" + values);
+ }
+ }));
+ PAssert.that(result).containsInAnyOrder("foo:[3, 3, 4, 5]");
+
+ List<String> messages = new CopyOnWriteArrayList<>();
+ JobState.Enum state = run(p, options, "bounded", messages);
+ assertEquals(String.join("\n", messages), JobState.Enum.DONE, state);
+ // The runner leaves the streaming option as submitted.
+ assertTrue(options.isStreaming());
+ }
+
+ private static SparkPipelineOptions options() {
+ PipelineOptions options =
PipelineOptionsFactory.fromArgs("--experiments=beam_fn_api").create();
+ options.setRunner(CrashingRunner.class);
+ options
+ .as(PortablePipelineOptions.class)
+ .setDefaultEnvironmentType(Environments.ENVIRONMENT_EMBEDDED);
+ SparkPipelineOptions sparkOptions = options.as(SparkPipelineOptions.class);
+ sparkOptions.setSparkMaster("local[2]");
+ sparkOptions.setStreaming(true);
+ sparkOptions.setUseStructuredStreaming(true);
+ return sparkOptions;
+ }
+
+ /** Submits the pipeline through the job invoker and returns its terminal
state. */
+ private static JobState.Enum run(
+ Pipeline p, SparkPipelineOptions options, String jobId, List<String>
messages)
+ throws Exception {
+ RunnerApi.Pipeline pipelineProto = PipelineTranslation.toProto(p);
+ JobInvocation invocation =
+ SparkJobInvoker.createJobInvocation(
+ jobId, "fakeRetrievalToken", executor, pipelineProto, options);
+ invocation.addMessageListener(message ->
messages.add(message.getMessageText()));
+ invocation.start();
+ while (!JobInvocation.isTerminated(invocation.getState())) {
+ Thread.sleep(200);
+ }
+ return invocation.getState();
+ }
+}
diff --git
a/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslatorTest.java
b/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslatorTest.java
new file mode 100644
index 00000000000..1be1a58b9ae
--- /dev/null
+++
b/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetPortablePipelineTranslatorTest.java
@@ -0,0 +1,442 @@
+/*
+ * 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.translation;
+
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getInputId;
+import static
org.apache.beam.runners.fnexecution.translation.PipelineTranslatorUtils.getOutputId;
+import static org.apache.beam.sdk.values.WindowedValues.valueInGlobalWindow;
+import static org.hamcrest.MatcherAssert.assertThat;
+import static org.hamcrest.Matchers.containsInAnyOrder;
+import static org.hamcrest.Matchers.containsString;
+import static org.junit.Assert.assertArrayEquals;
+import static org.junit.Assert.assertEquals;
+import static org.junit.Assert.assertThrows;
+
+import java.io.Serializable;
+import java.util.ArrayList;
+import java.util.Arrays;
+import java.util.Collections;
+import java.util.List;
+import java.util.Map;
+import java.util.stream.Collectors;
+import org.apache.beam.model.pipeline.v1.RunnerApi;
+import org.apache.beam.model.pipeline.v1.RunnerApi.ExecutableStagePayload;
+import org.apache.beam.runners.fnexecution.provisioning.JobInfo;
+import org.apache.beam.runners.spark.SparkContextRule;
+import org.apache.beam.runners.spark.SparkPipelineOptions;
+import org.apache.beam.runners.spark.metrics.MetricsAccumulator;
+import org.apache.beam.sdk.Pipeline;
+import org.apache.beam.sdk.coders.KvCoder;
+import org.apache.beam.sdk.coders.NullableCoder;
+import org.apache.beam.sdk.coders.StringUtf8Coder;
+import org.apache.beam.sdk.coders.VarLongCoder;
+import org.apache.beam.sdk.options.PipelineOptionsFactory;
+import org.apache.beam.sdk.options.PortablePipelineOptions;
+import org.apache.beam.sdk.transforms.DoFn;
+import org.apache.beam.sdk.transforms.Flatten;
+import org.apache.beam.sdk.transforms.GroupByKey;
+import org.apache.beam.sdk.transforms.Impulse;
+import org.apache.beam.sdk.transforms.ParDo;
+import org.apache.beam.sdk.transforms.Reshuffle;
+import org.apache.beam.sdk.transforms.View;
+import org.apache.beam.sdk.transforms.windowing.FixedWindows;
+import org.apache.beam.sdk.transforms.windowing.GlobalWindow;
+import org.apache.beam.sdk.transforms.windowing.IntervalWindow;
+import org.apache.beam.sdk.transforms.windowing.PaneInfo;
+import org.apache.beam.sdk.transforms.windowing.Window;
+import org.apache.beam.sdk.util.construction.Environments;
+import org.apache.beam.sdk.util.construction.PipelineOptionsTranslation;
+import org.apache.beam.sdk.util.construction.PipelineTranslation;
+import org.apache.beam.sdk.util.construction.graph.ExecutableStage;
+import org.apache.beam.sdk.util.construction.graph.GreedyPipelineFuser;
+import org.apache.beam.sdk.util.construction.graph.PipelineNode;
+import org.apache.beam.sdk.util.construction.graph.PipelineNode.PTransformNode;
+import
org.apache.beam.sdk.util.construction.graph.TrivialNativeTransformExpander;
+import org.apache.beam.sdk.values.KV;
+import org.apache.beam.sdk.values.PCollection;
+import org.apache.beam.sdk.values.PCollectionList;
+import org.apache.beam.sdk.values.PCollectionTuple;
+import org.apache.beam.sdk.values.PCollectionView;
+import org.apache.beam.sdk.values.TupleTag;
+import org.apache.beam.sdk.values.TupleTagList;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.beam.sdk.values.WindowedValues;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Iterables;
+import org.apache.spark.sql.Dataset;
+import org.joda.time.Duration;
+import org.joda.time.Instant;
+import org.junit.After;
+import org.junit.Before;
+import org.junit.ClassRule;
+import org.junit.Test;
+import org.junit.runner.RunWith;
+import org.junit.runners.JUnit4;
+
+/**
+ * Unit tests for {@link SparkDatasetPortablePipelineTranslator}. Runner-side
transforms are
+ * translated on injected Datasets. Executable stages run in the embedded SDK
harness.
+ */
+@RunWith(JUnit4.class)
+public class SparkDatasetPortablePipelineTranslatorTest implements
Serializable {
+
+ @ClassRule public static SparkContextRule contextRule = new
SparkContextRule("local[2]");
+
+ private transient SparkPipelineOptions options;
+ private transient SparkDatasetPortablePipelineTranslator translator;
+ private transient SparkDatasetTranslationContext context;
+
+ @Before
+ public void setUp() {
+ options = PipelineOptionsFactory.create().as(SparkPipelineOptions.class);
+ options.setUseStructuredStreaming(true);
+ options
+ .as(PortablePipelineOptions.class)
+ .setDefaultEnvironmentType(Environments.ENVIRONMENT_EMBEDDED);
+ MetricsAccumulator.clear();
+ MetricsAccumulator.init(options, contextRule.getSparkContext());
+ translator = new SparkDatasetPortablePipelineTranslator();
+ context =
+ translator.createTranslationContext(
+ contextRule.getSparkContext(),
+ options,
+ JobInfo.create("job", "job", "token",
PipelineOptionsTranslation.toProto(options)));
+ }
+
+ @After
+ public void tearDown() {
+ context.getSparkSession().sharedState().cacheManager().clearCache();
+ MetricsAccumulator.clear();
+ }
+
+ @Test
+ public void impulseIsOneEmptyElementInTheGlobalWindow() {
+ Pipeline p = Pipeline.create(options);
+ p.apply("impulse", Impulse.create());
+ RunnerApi.Pipeline pipeline = PipelineTranslation.toProto(p);
+
+ translator.translate(pipeline, context);
+
+ List<WindowedValue<byte[]>> result =
collect(getOutputId(transformNamed(pipeline, "impulse")));
+ assertEquals(1, result.size());
+ assertArrayEquals(new byte[0], result.get(0).getValue());
+ assertEquals(
+ Collections.singletonList(GlobalWindow.INSTANCE),
+ new ArrayList<>(result.get(0).getWindows()));
+ }
+
+ @Test
+ public void flattenReEncodesInputsWithAnotherCoder() {
+ Pipeline p = Pipeline.create(options);
+ PCollection<KV<String, Long>> plain =
+ p.apply("impulse", Impulse.create())
+ .apply("plain", ParDo.of(new Placeholder<KV<String, Long>>()))
+ .setCoder(KvCoder.of(StringUtf8Coder.of(), VarLongCoder.of()));
+ // Same Java type as the first input, different encoding.
+ PCollection<KV<String, Long>> nullable =
+ p.apply("impulse2", Impulse.create())
+ .apply("nullable", ParDo.of(new Placeholder<KV<String, Long>>()))
+ .setCoder(KvCoder.of(NullableCoder.of(StringUtf8Coder.of()),
VarLongCoder.of()));
+ PCollectionList.of(plain).and(nullable).apply("flatten",
Flatten.pCollections());
+ RunnerApi.Pipeline pipeline = PipelineTranslation.toProto(p);
+ PTransformNode flatten = transformNamed(pipeline, "flatten");
+ inject(
+ getOutputId(transformNamed(pipeline, "plain")),
+ pipeline,
+ Arrays.asList(valueInGlobalWindow(KV.of("a", 1L)),
valueInGlobalWindow(KV.of("b", 2L))));
+ inject(
+ getOutputId(transformNamed(pipeline, "nullable")),
+ pipeline,
+ Collections.singletonList(valueInGlobalWindow(KV.of("c", 3L))));
+
+ SparkDatasetPortablePipelineTranslator.translateFlatten(flatten, pipeline,
context);
+
+ assertThat(
+ values(collect(getOutputId(flatten))),
+ containsInAnyOrder(KV.of("a", 1L), KV.of("b", 2L), KV.of("c", 3L)));
+ }
+
+ @Test
+ public void groupByKeyGroupsPerKeyAndWindow() {
+ Pipeline p = Pipeline.create(options);
+ p.apply("impulse", Impulse.create())
+ .apply("kv", ParDo.of(new Placeholder<KV<String, Long>>()))
+ .setCoder(KvCoder.of(StringUtf8Coder.of(), VarLongCoder.of()))
+ .apply("window",
Window.into(FixedWindows.of(Duration.standardSeconds(10))))
+ .apply("gbk", GroupByKey.create());
+ RunnerApi.Pipeline pipeline = PipelineTranslation.toProto(p);
+ PTransformNode gbk = transformNamed(pipeline, "gbk");
+ IntervalWindow first = new IntervalWindow(new Instant(0),
Duration.standardSeconds(10));
+ IntervalWindow second = new IntervalWindow(new Instant(10_000),
Duration.standardSeconds(10));
+ inject(
+ getInputId(gbk),
+ pipeline,
+ Arrays.asList(
+ WindowedValues.of(KV.of("a", 1L), new Instant(1_000), first,
PaneInfo.NO_FIRING),
+ WindowedValues.of(KV.of("a", 2L), new Instant(2_000), first,
PaneInfo.NO_FIRING),
+ WindowedValues.of(KV.of("b", 3L), new Instant(3_000), first,
PaneInfo.NO_FIRING),
+ WindowedValues.of(KV.of("a", 4L), new Instant(14_000), second,
PaneInfo.NO_FIRING)));
+
+ SparkDatasetPortablePipelineTranslator.translateGroupByKey(gbk, pipeline,
context);
+
+ List<String> groups = new ArrayList<>();
+ for (WindowedValue<KV<String, Iterable<Long>>> group :
+ this.<KV<String, Iterable<Long>>>collect(getOutputId(gbk))) {
+ List<Long> values = new ArrayList<>();
+ group.getValue().getValue().forEach(values::add);
+ Collections.sort(values);
+ groups.add(
+ group.getValue().getKey()
+ + values
+ + Iterables.getOnlyElement(group.getWindows())
+ + "@"
+ + group.getTimestamp());
+ }
+ assertThat(
+ groups,
+ containsInAnyOrder(
+ "a[1, 2]" + first + "@" + first.maxTimestamp(),
+ "b[3]" + first + "@" + first.maxTimestamp(),
+ "a[4]" + second + "@" + second.maxTimestamp()));
+ }
+
+ @Test
+ public void reshuffleRepartitionsAndKeepsEveryElement() {
+ Pipeline p = Pipeline.create(options);
+ p.apply("impulse", Impulse.create())
+ .apply("kv", ParDo.of(new Placeholder<KV<String, Long>>()))
+ .setCoder(KvCoder.of(StringUtf8Coder.of(), VarLongCoder.of()))
+ .apply("reshuffle", Reshuffle.of());
+ RunnerApi.Pipeline pipeline = PipelineTranslation.toProto(p);
+ PTransformNode reshuffle = transformNamed(pipeline, "reshuffle");
+ inject(
+ getInputId(reshuffle),
+ pipeline,
+ Arrays.asList(
+ valueInGlobalWindow(KV.of("a", 1L)),
+ valueInGlobalWindow(KV.of("a", 2L)),
+ valueInGlobalWindow(KV.of("b", 3L))));
+
+ SparkDatasetPortablePipelineTranslator.translateReshuffle(reshuffle,
pipeline, context);
+
+ Dataset<WindowedValue<KV<String, Long>>> output =
context.getDataset(getOutputId(reshuffle));
+ assertEquals(
+ contextRule.getSparkContext().defaultParallelism().intValue(),
+ output.rdd().getNumPartitions());
+ assertThat(
+ values(output.collectAsList()),
+ containsInAnyOrder(KV.of("a", 1L), KV.of("a", 2L), KV.of("b", 3L)));
+ }
+
+ @Test
+ public void executableStageOutputsAreDemultiplexedPerTag() {
+ TupleTag<KV<String, String>> words = new TupleTag<KV<String,
String>>("words") {};
+ TupleTag<KV<String, Long>> lengths = new TupleTag<KV<String,
Long>>("lengths") {};
+ Pipeline p = Pipeline.create(options);
+ PCollectionTuple outputs =
+ p.apply("impulse", Impulse.create())
+ .apply(
+ "split",
+ ParDo.of(
+ new DoFn<byte[], KV<String, String>>() {
+ @ProcessElement
+ public void process(MultiOutputReceiver out) {
+ for (String word : Arrays.asList("one", "three")) {
+ out.get(words).output(KV.of("word", word));
+ out.get(lengths).output(KV.of("length", (long)
word.length()));
+ }
+ }
+ })
+ .withOutputTags(words, TupleTagList.of(lengths)));
+ // A stage only emits outputs that a runner-side transform reads.
GroupByKey is one.
+ outputs.get(words).apply("groupWords", GroupByKey.create());
+ outputs.get(lengths).apply("groupLengths", GroupByKey.create());
+ RunnerApi.Pipeline pipeline = fused(p);
+ assertEquals(2, onlyStage(pipeline).getOutputsCount());
+
+ translator.translate(pipeline, context);
+
+ Map<String, String> split = transformNamed(pipeline,
"split").getTransform().getOutputsMap();
+ assertThat(
+ values(collect(split.get(words.getId()))),
+ containsInAnyOrder(KV.of("word", "one"), KV.of("word", "three")));
+ assertThat(
+ values(collect(split.get(lengths.getId()))),
+ containsInAnyOrder(KV.of("length", 3L), KV.of("length", 5L)));
+ }
+
+ @Test
+ public void sideInputsAreBroadcastToTheStage() {
+ Pipeline p = Pipeline.create(options);
+ PCollectionView<Iterable<String>> view =
+ p.apply("impulse", Impulse.create())
+ .apply(
+ "words",
+ ParDo.of(
+ new DoFn<byte[], String>() {
+ @ProcessElement
+ public void process(OutputReceiver<String> out) {
+ out.output("one");
+ out.output("three");
+ }
+ }))
+ .apply("view", View.asIterable());
+ p.apply("impulse2", Impulse.create())
+ .apply(
+ "total",
+ ParDo.of(
+ new DoFn<byte[], KV<String, Long>>() {
+ @ProcessElement
+ public void process(ProcessContext c) {
+ long total = 0;
+ for (String word : c.sideInput(view)) {
+ total += word.length();
+ }
+ c.output(KV.of("total", total));
+ }
+ })
+ .withSideInputs(view))
+ .apply("groupTotal", GroupByKey.create());
+ RunnerApi.Pipeline pipeline = fused(p);
+
+ translator.translate(pipeline, context);
+
+ assertThat(
+ values(collect(getOutputId(transformNamed(pipeline, "total")))),
+ containsInAnyOrder(KV.of("total", 8L)));
+ }
+
+ @Test
+ public void unboundedInputIsRejected() {
+ RunnerApi.Pipeline pipeline =
+ RunnerApi.Pipeline.newBuilder()
+ .setComponents(
+ RunnerApi.Components.newBuilder()
+ .putPcollections(
+ "unbounded",
+ RunnerApi.PCollection.newBuilder()
+ .setIsBounded(RunnerApi.IsBounded.Enum.UNBOUNDED)
+ .build()))
+ .build();
+
+ UnsupportedOperationException thrown =
+ assertThrows(
+ UnsupportedOperationException.class, () ->
translator.translate(pipeline, context));
+ assertThat(thrown.getMessage(), containsString("unbounded input"));
+ }
+
+ @Test
+ public void stageWithUserStateIsRejected() {
+ ExecutableStagePayload payload =
+ ExecutableStagePayload.newBuilder()
+ .setInput("input")
+ .addUserStates(
+ ExecutableStagePayload.UserStateId.newBuilder()
+ .setTransformId("pardo")
+ .setLocalName("count"))
+ .build();
+
+ UnsupportedOperationException thrown =
+ assertThrows(
+ UnsupportedOperationException.class,
+ () ->
+
SparkDatasetPortablePipelineTranslator.translateExecutableStage(
+ stage(payload), RunnerApi.Pipeline.getDefaultInstance(),
context));
+ assertThat(thrown.getMessage(), containsString("state or timers"));
+ }
+
+ @Test
+ public void stageWithTimersIsRejected() {
+ ExecutableStagePayload payload =
+ ExecutableStagePayload.newBuilder()
+ .setInput("input")
+ .addTimers(
+ ExecutableStagePayload.TimerId.newBuilder()
+ .setTransformId("pardo")
+ .setLocalName("expiry"))
+ .build();
+
+ UnsupportedOperationException thrown =
+ assertThrows(
+ UnsupportedOperationException.class,
+ () ->
+
SparkDatasetPortablePipelineTranslator.translateExecutableStage(
+ stage(payload), RunnerApi.Pipeline.getDefaultInstance(),
context));
+ assertThat(thrown.getMessage(), containsString("state or timers"));
+ }
+
+ /** A DoFn that only gives its output PCollection a coder. It is never
translated or run. */
+ private static class Placeholder<T> extends DoFn<byte[], T> {
+ @ProcessElement
+ public void process() {}
+ }
+
+ private RunnerApi.Pipeline fused(Pipeline p) {
+ RunnerApi.Pipeline pipeline =
+ TrivialNativeTransformExpander.forKnownUrns(
+ PipelineTranslation.toProto(p), translator.knownUrns());
+ return GreedyPipelineFuser.fuse(pipeline).toPipeline();
+ }
+
+ private static PTransformNode transformNamed(RunnerApi.Pipeline pipeline,
String name) {
+ for (Map.Entry<String, RunnerApi.PTransform> transform :
+ pipeline.getComponents().getTransformsMap().entrySet()) {
+ if (name.equals(transform.getValue().getUniqueName())) {
+ return PipelineNode.pTransform(transform.getKey(),
transform.getValue());
+ }
+ }
+ throw new IllegalArgumentException("No transform named " + name);
+ }
+
+ private static RunnerApi.PTransform onlyStage(RunnerApi.Pipeline pipeline) {
+ return Iterables.getOnlyElement(
+ pipeline.getComponents().getTransformsMap().values().stream()
+ .filter(transform ->
ExecutableStage.URN.equals(transform.getSpec().getUrn()))
+ .collect(Collectors.toList()));
+ }
+
+ private static PTransformNode stage(ExecutableStagePayload payload) {
+ return PipelineNode.pTransform(
+ "stage",
+ RunnerApi.PTransform.newBuilder()
+ .putInputs("input", payload.getInput())
+ .setSpec(
+ RunnerApi.FunctionSpec.newBuilder()
+ .setUrn(ExecutableStage.URN)
+ .setPayload(payload.toByteString()))
+ .build());
+ }
+
+ /** Registers {@code values} as the Dataset of a PCollection, in a single
partition. */
+ private <T> void inject(
+ String pCollectionId, RunnerApi.Pipeline pipeline,
List<WindowedValue<T>> values) {
+ Dataset<WindowedValue<T>> dataset =
+ context
+ .getSparkSession()
+ .createDataset(values, context.windowedEncoder(pCollectionId,
pipeline.getComponents()))
+ .coalesce(1);
+ context.putDataset(pCollectionId, dataset);
+ }
+
+ private <T> List<WindowedValue<T>> collect(String pCollectionId) {
+ return context.<T>getDataset(pCollectionId).collectAsList();
+ }
+
+ private static <T> List<T> values(List<WindowedValue<T>> windowedValues) {
+ return
windowedValues.stream().map(WindowedValue::getValue).collect(Collectors.toList());
+ }
+}
diff --git
a/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContextTest.java
b/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContextTest.java
new file mode 100644
index 00000000000..74b78f24ecd
--- /dev/null
+++
b/runners/spark/src/test/java/org/apache/beam/runners/spark/translation/SparkDatasetTranslationContextTest.java
@@ -0,0 +1,164 @@
+/*
+ * 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.translation;
+
+import static org.apache.beam.sdk.values.WindowedValues.valueInGlobalWindow;
+import static org.junit.Assert.assertEquals;
+
+import java.nio.charset.StandardCharsets;
+import java.util.Collections;
+import java.util.Set;
+import java.util.UUID;
+import java.util.concurrent.ConcurrentHashMap;
+import org.apache.beam.runners.fnexecution.provisioning.JobInfo;
+import org.apache.beam.runners.spark.SparkContextRule;
+import org.apache.beam.runners.spark.SparkPipelineOptions;
+import
org.apache.beam.runners.spark.structuredstreaming.translation.helpers.EncoderHelpers;
+import org.apache.beam.sdk.coders.ByteArrayCoder;
+import org.apache.beam.sdk.options.PipelineOptionsFactory;
+import org.apache.beam.sdk.transforms.windowing.GlobalWindow;
+import org.apache.beam.sdk.util.construction.PipelineOptionsTranslation;
+import org.apache.beam.sdk.values.WindowedValue;
+import org.apache.spark.api.java.function.MapFunction;
+import org.apache.spark.sql.Dataset;
+import org.apache.spark.sql.Encoder;
+import org.apache.spark.storage.StorageLevel;
+import org.junit.After;
+import org.junit.ClassRule;
+import org.junit.Test;
+import org.junit.runner.RunWith;
+import org.junit.runners.JUnit4;
+
+/** Unit tests for {@link SparkDatasetTranslationContext}. */
+@RunWith(JUnit4.class)
+public class SparkDatasetTranslationContextTest {
+
+ @ClassRule public static SparkContextRule contextRule = new
SparkContextRule();
+
+ private static final String COLLECTION = "collection";
+
+ /** Names of the Datasets that were evaluated, recorded from the executor
side. */
+ private static final Set<String> EVALUATED = ConcurrentHashMap.newKeySet();
+
+ private SparkDatasetTranslationContext context;
+
+ @After
+ public void tearDown() {
+ if (context != null) {
+ context.getSparkSession().sharedState().cacheManager().clearCache();
+ }
+ }
+
+ @Test
+ public void datasetReadBySeveralTransformsIsPersisted() {
+ SparkPipelineOptions options = options();
+ context = newContext(options);
+ context.addConsumer(COLLECTION);
+ context.addConsumer(COLLECTION);
+
+ context.putDataset(COLLECTION, dataset());
+
+ assertEquals(
+ StorageLevel.fromString(options.getStorageLevel()),
+ context.getDataset(COLLECTION).storageLevel());
+ }
+
+ @Test
+ public void datasetReadByOneTransformIsNotPersisted() {
+ context = newContext(options());
+ context.addConsumer(COLLECTION);
+
+ context.putDataset(COLLECTION, dataset());
+
+ assertEquals(StorageLevel.NONE(),
context.getDataset(COLLECTION).storageLevel());
+ }
+
+ @Test
+ public void projectionOfPersistedRowsIsNotPersistedAgain() {
+ context = newContext(options());
+ context.addConsumer(COLLECTION);
+ context.addConsumer(COLLECTION);
+
+ context.putDataset(COLLECTION, dataset(), false);
+
+ assertEquals(StorageLevel.NONE(),
context.getDataset(COLLECTION).storageLevel());
+ }
+
+ @Test
+ public void cacheDisabledSkipsPersisting() {
+ SparkPipelineOptions options = options();
+ options.setCacheDisabled(true);
+ context = newContext(options);
+ context.addConsumer(COLLECTION);
+ context.addConsumer(COLLECTION);
+
+ context.putDataset(COLLECTION, dataset());
+
+ assertEquals(StorageLevel.NONE(),
context.getDataset(COLLECTION).storageLevel());
+ }
+
+ @Test
+ public void computeOutputsEvaluatesOnlyUnconsumedDatasets() {
+ EVALUATED.clear();
+ context = newContext(options());
+ context.putDataset("leaf", recording("leaf"));
+ context.putDataset("consumed", recording("consumed"));
+ context.getDataset("consumed");
+
+ context.computeOutputs();
+
+ assertEquals(Collections.singleton("leaf"), EVALUATED);
+ }
+
+ private static SparkPipelineOptions options() {
+ return PipelineOptionsFactory.create().as(SparkPipelineOptions.class);
+ }
+
+ private static SparkDatasetTranslationContext
newContext(SparkPipelineOptions options) {
+ return new SparkDatasetTranslationContext(
+ contextRule.getSparkContext(),
+ options,
+ JobInfo.create("job", "job", "token",
PipelineOptionsTranslation.toProto(options)));
+ }
+
+ private static Encoder<WindowedValue<byte[]>> encoder() {
+ return EncoderHelpers.windowedValueEncoder(
+ EncoderHelpers.encoderFor(ByteArrayCoder.of()),
+ EncoderHelpers.encoderFor(GlobalWindow.Coder.INSTANCE));
+ }
+
+ /** A one element Dataset with a unique payload, so no two tests share a
cached plan. */
+ private Dataset<WindowedValue<byte[]>> dataset() {
+ byte[] payload =
UUID.randomUUID().toString().getBytes(StandardCharsets.UTF_8);
+ return context
+ .getSparkSession()
+
.createDataset(Collections.singletonList(valueInGlobalWindow(payload)),
encoder());
+ }
+
+ /** A Dataset that records {@code name} in {@link #EVALUATED} when its rows
are computed. */
+ private Dataset<WindowedValue<byte[]>> recording(String name) {
+ return dataset()
+ .map(
+ (MapFunction<WindowedValue<byte[]>, WindowedValue<byte[]>>)
+ value -> {
+ EVALUATED.add(name);
+ return value;
+ },
+ encoder());
+ }
+}