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 000000000000..4c6e3309e306 --- /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 5d2d832f3003..ff94d5d78b42 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`| [![.github/workflows/beam_PostCommit_Java_PVR_Spark_Batch.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark_Batch.yml/badge.svg?event=schedule)](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`| [![.github/workflows/beam_PostCommit_Java_PVR_Spark4_Batch.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Batch.yml/badge.svg?event=schedule)](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`| [![.github/workflows/beam_PostCommit_Java_PVR_Spark4_Streaming.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Streaming.yml/badge.svg?event=schedule)](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_Streaming.yml?query=event%3Aschedule) | +| [ 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`| [![.github/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml/badge.svg?event=schedule)](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_PVR_Spark4_StructuredStreaming.yml?query=event%3Aschedule) | | [ 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`| [![.github/workflows/beam_PostCommit_Java_Tpcds_Dataflow.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Dataflow.yml/badge.svg?event=schedule)](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`| [![.github/workflows/beam_PostCommit_Java_Tpcds_Flink.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Flink.yml/badge.svg?event=schedule)](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`| [![.github/workflows/beam_PostCommit_Java_Tpcds_Spark.yml](https://github.com/apache/beam/actions/workflows/beam_PostCommit_Java_Tpcds_Spark.yml/badge.svg?event=schedule)](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 000000000000..33c8a3be79e9 --- /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 371a44512657..02ba1f0b5e60 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 0166f0654018..637ac98375da 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 sickbayTests) { +def portableValidatesRunnerTask(String name, boolean streaming, boolean docker, ArrayList 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 2ad149431d21..f2d5b2c1272c 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 91a94896b89b..51fe26d889fd 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.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 SparkPipelineRunner(SparkPipelineOptions pipelineOptions) { @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 PortablePipelineResult run(RunnerApi.Pipeline pipeline, JobInfo jobInfo) 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 PortablePipelineResult run(RunnerApi.Pipeline pipeline, JobInfo jobInfo) 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 612d71b1aea1..0e06b5f66e30 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 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 000000000000..b58920ab198b --- /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. + * + *

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 { + + private final ImmutableMap urnToTransformTranslator; + + interface PTransformTranslator { + void translate( + PTransformNode transformNode, + RunnerApi.Pipeline pipeline, + SparkDatasetTranslationContext context); + } + + public SparkDatasetPortablePipelineTranslator() { + ImmutableMap.Builder 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 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> dataset = + context + .getSparkSession() + .createDataset( + Collections.singletonList(WindowedValues.valueInGlobalWindow(new byte[0])), + context.windowedEncoder(outputId, pipeline.getComponents())); + context.putDataset(outputId, dataset); + } + + @VisibleForTesting + static 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> input = context.getDataset(inputId); + Map outputs = transformNode.getTransform().getOutputsMap(); + BiMap outputMap = createOutputMap(outputs.values()); + Coder windowCoder = getWindowingStrategy(inputId, components).getWindowFn().windowCoder(); + + SparkExecutableStageFunction 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> sink = + input.mapPartitions( + (MapPartitionsFunction, WindowedValue>) + elements -> { + Iterator 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>> encoders = + new ArrayList<>(Collections.nCopies(outputMap.size(), null)); + for (Map.Entry output : outputMap.entrySet()) { + encoders.set(output.getValue(), context.windowedEncoder(output.getKey(), components)); + } + Dataset>> staged = + input.mapPartitions( + (MapPartitionsFunction, Tuple2>>) + 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 output : outputMap.entrySet()) { + int tag = output.getValue(); + TypedColumn>, WindowedValue> 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>> tagged( + Iterator values) { + return Iterators.transform( + values, + value -> new Tuple2<>(value.getUnionTag(), (WindowedValue) value.getValue())); + } + + /** Collects each side input of a stage and broadcasts its encoded elements. */ + private static + ImmutableMap>, WindowedValueCoder>> + broadcastSideInputs( + RunnerApi.ExecutableStagePayload stagePayload, + SparkDatasetTranslationContext context) { + Map>, WindowedValueCoder>> 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 coder = getWindowedValueCoder(collectionId, components); + Dataset> dataset = context.getDataset(collectionId); + List bytes = + new ArrayList<>( + dataset + .map( + (MapFunction, 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 void translateGroupByKey( + PTransformNode transformNode, + RunnerApi.Pipeline pipeline, + SparkDatasetTranslationContext context) { + RunnerApi.Components components = pipeline.getComponents(); + String inputId = getInputId(transformNode); + String outputId = getOutputId(transformNode); + Dataset>> input = context.getDataset(inputId); + WindowedValueCoder> inputCoder = getWindowedValueCoder(inputId, components); + KvCoder kvCoder = (KvCoder) inputCoder.getValueCoder(); + Coder keyCoder = kvCoder.getKeyCoder(); + WindowingStrategy windowingStrategy = + getWindowingStrategy(inputId, components); + + // Batch semantics: all values of a key are present, so every window of the key can close. + SparkGroupAlsoByWindowViaOutputBufferFn groupAlsoByWindow = + new SparkGroupAlsoByWindowViaOutputBufferFn<>( + windowingStrategy, + new TranslationUtils.InMemoryStateInternalsFactory<>(), + SystemReduceFn.buffering(kvCoder.getValueCoder()), + context.getSerializableOptions()); + + Dataset>>> grouped = + input + .groupByKey( + (MapFunction>, byte[]>) + value -> CoderHelpers.toByteArray(value.getValue().getKey(), keyCoder), + Encoders.BINARY()) + .flatMapGroups( + (FlatMapGroupsFunction< + byte[], WindowedValue>, WindowedValue>>>) + (keyBytes, values) -> { + K key = CoderHelpers.fromByteArray(keyBytes, keyCoder); + List> windowedValues = new ArrayList<>(); + while (values.hasNext()) { + WindowedValue> value = values.next(); + windowedValues.add(value.withValue(value.getValue().getValue())); + } + return groupAlsoByWindow.call( + KV.>>of(key, windowedValues)); + }, + context.windowedEncoder(outputId, components)); + context.putDataset(outputId, grouped); + } + + @VisibleForTesting + static void translateFlatten( + PTransformNode transformNode, + RunnerApi.Pipeline pipeline, + SparkDatasetTranslationContext context) { + RunnerApi.Components components = pipeline.getComponents(); + String outputId = getOutputId(transformNode); + WindowedValueCoder outputCoder = getWindowedValueCoder(outputId, components); + Encoder> outputEncoder = context.windowedEncoder(outputId, components); + Dataset> result = null; + for (String inputId : transformNode.getTransform().getInputsMap().values()) { + Dataset> 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>) v -> v, outputEncoder); + } + result = result == null ? input : result.union(input); + } + if (result == null) { + result = context.getSparkSession().emptyDataset(outputEncoder); + } + context.putDataset(outputId, result); + } + + @VisibleForTesting + static void translateReshuffle( + PTransformNode transformNode, + RunnerApi.Pipeline pipeline, + SparkDatasetTranslationContext context) { + Dataset> 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 000000000000..d58fd62cd7e2 --- /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 consumers = new HashMap<>(); + private final Map datasets = new HashMap<>(); + private final Set leaves = new LinkedHashSet<>(); + private final Map, 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 void putDataset(String pCollectionId, Dataset> 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 void putDataset( + String pCollectionId, Dataset> 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 Dataset> getDataset(String pCollectionId) { + leaves.remove(pCollectionId); + return datasets.get(pCollectionId); + } + + /** Encoder of the windowed values of a PCollection, derived from its wire coder. */ + public Encoder> windowedEncoder( + String pCollectionId, RunnerApi.Components components) { + Coder valueCoder = + PipelineTranslatorUtils.getWindowedValueCoder(pCollectionId, components).getValueCoder(); + Coder windowCoder = + PipelineTranslatorUtils.getWindowingStrategy(pCollectionId, components) + .getWindowFn() + .windowCoder(); + return EncoderHelpers.windowedValueEncoder(encoderOf(valueCoder), encoderOf(windowCoder)); + } + + private Encoder encoderOf(Coder coder) { + return (Encoder) 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 000000000000..90375d7f6e24 --- /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 words = + p.apply("impulse", Impulse.create()) + .apply( + "create", + ParDo.of( + new DoFn() { + @ProcessElement + public void process(ProcessContext ctxt) { + ctxt.output("zero"); + ctxt.output("one"); + ctxt.output("two"); + } + })); + PCollection more = + p.apply("impulse2", Impulse.create()) + .apply( + "create2", + ParDo.of( + new DoFn() { + @ProcessElement + public void process(ProcessContext ctxt) { + ctxt.output("three"); + } + })); + PCollection result = + PCollectionList.of(words) + .and(more) + .apply("flatten", Flatten.pCollections()) + .apply( + "len", + ParDo.of( + new DoFn() { + @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>, String>() { + @ProcessElement + public void process(ProcessContext ctxt) { + // The order of grouped values is not defined, so sort before comparing. + List 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 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 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 000000000000..1be1a58b9aed --- /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> 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> plain = + p.apply("impulse", Impulse.create()) + .apply("plain", ParDo.of(new Placeholder>())) + .setCoder(KvCoder.of(StringUtf8Coder.of(), VarLongCoder.of())); + // Same Java type as the first input, different encoding. + PCollection> nullable = + p.apply("impulse2", Impulse.create()) + .apply("nullable", ParDo.of(new Placeholder>())) + .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>())) + .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 groups = new ArrayList<>(); + for (WindowedValue>> group : + this.>>collect(getOutputId(gbk))) { + List 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>())) + .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>> 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> words = new TupleTag>("words") {}; + TupleTag> lengths = new TupleTag>("lengths") {}; + Pipeline p = Pipeline.create(options); + PCollectionTuple outputs = + p.apply("impulse", Impulse.create()) + .apply( + "split", + ParDo.of( + new DoFn>() { + @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 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> view = + p.apply("impulse", Impulse.create()) + .apply( + "words", + ParDo.of( + new DoFn() { + @ProcessElement + public void process(OutputReceiver out) { + out.output("one"); + out.output("three"); + } + })) + .apply("view", View.asIterable()); + p.apply("impulse2", Impulse.create()) + .apply( + "total", + ParDo.of( + new DoFn>() { + @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 extends DoFn { + @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 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 void inject( + String pCollectionId, RunnerApi.Pipeline pipeline, List> values) { + Dataset> dataset = + context + .getSparkSession() + .createDataset(values, context.windowedEncoder(pCollectionId, pipeline.getComponents())) + .coalesce(1); + context.putDataset(pCollectionId, dataset); + } + + private List> collect(String pCollectionId) { + return context.getDataset(pCollectionId).collectAsList(); + } + + private static List values(List> 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 000000000000..74b78f24ecd0 --- /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 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> 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> 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> recording(String name) { + return dataset() + .map( + (MapFunction, WindowedValue>) + value -> { + EVALUATED.add(name); + return value; + }, + encoder()); + } +}