diff --git a/scio-smb/src/main/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransform.java b/scio-smb/src/main/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransform.java index 44c44455c3..b4b0a2457e 100644 --- a/scio-smb/src/main/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransform.java +++ b/scio-smb/src/main/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransform.java @@ -53,12 +53,14 @@ import org.apache.beam.sdk.transforms.Filter; import org.apache.beam.sdk.transforms.PTransform; import org.apache.beam.sdk.transforms.ParDo; +import org.apache.beam.sdk.transforms.Reshuffle; import org.apache.beam.sdk.transforms.display.DisplayData; import org.apache.beam.sdk.transforms.join.CoGbkResult; import org.apache.beam.sdk.transforms.join.CoGbkResultSchema; import org.apache.beam.sdk.transforms.windowing.BoundedWindow; import org.apache.beam.sdk.values.KV; import org.apache.beam.sdk.values.PBegin; +import org.apache.beam.sdk.values.PCollection; import org.apache.beam.sdk.values.PCollectionView; import org.apache.beam.sdk.values.TupleTag; import org.apache.beam.sdk.values.TupleTagList; @@ -83,6 +85,7 @@ public class SortedBucketTransform extends PTransform bucketSource; private final DoFn, KV> finalizeBuckets; private final ParDo.SingleOutput doFn; + private final boolean usesSideInputs; public SortedBucketTransform( List> sources, @@ -99,13 +102,14 @@ public SortedBucketTransform( String filenameSuffix, String filenamePrefix) { Preconditions.checkNotNull(outputDirectory, "outputDirectory is not set"); + usesSideInputs = sideInputTransformFn != null; Preconditions.checkState( - !((transformFn == null) && (sideInputTransformFn == null)), // at least one defined + transformFn != null || usesSideInputs, // at least one defined "At least one of transformFn and sideInputTransformFn must be set"); Preconditions.checkState( - !((transformFn != null) && (sideInputTransformFn != null)), // only one defined + transformFn == null || !usesSideInputs, // only one defined "At most one of transformFn and sideInputTransformFn may be set"); - if (sideInputTransformFn != null) { + if (usesSideInputs) { Preconditions.checkNotNull(sides, "If using sideInputTransformFn, sides must not be null"); } @@ -151,11 +155,18 @@ public SortedBucketTransform( @Override public final WriteResult expand(final PBegin begin) { - return WriteResult.fromTuple( + PCollection bucketOffsets = begin .getPipeline() // outputs bucket offsets for the various SMB readers - .apply("BucketOffsets", Read.from(bucketSource)) + .apply("BucketOffsets", Read.from(bucketSource)); + + if (usesSideInputs) { + bucketOffsets = bucketOffsets.apply("RedistributeBuckets", Reshuffle.viaRandomKey()); + } + + return WriteResult.fromTuple( + bucketOffsets .apply("MergeBuckets", this.doFn) .apply(Filter.by(Objects::nonNull)) .apply(Group.globally()) diff --git a/scio-smb/src/test/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransformTest.java b/scio-smb/src/test/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransformTest.java index ad6f8c1c33..7dc5e3bab2 100644 --- a/scio-smb/src/test/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransformTest.java +++ b/scio-smb/src/test/java/org/apache/beam/sdk/extensions/smb/SortedBucketTransformTest.java @@ -26,11 +26,14 @@ import java.util.Collections; import java.util.Comparator; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Set; import java.util.function.Function; import java.util.stream.Collectors; +import org.apache.beam.sdk.Pipeline; +import org.apache.beam.sdk.Pipeline.PipelineVisitor.CompositeBehavior; import org.apache.beam.sdk.PipelineResult; import org.apache.beam.sdk.coders.StringUtf8Coder; import org.apache.beam.sdk.extensions.smb.SMBFilenamePolicy.FileAssignment; @@ -41,6 +44,7 @@ import org.apache.beam.sdk.io.fs.MatchResult.Status; import org.apache.beam.sdk.io.fs.ResourceId; import org.apache.beam.sdk.metrics.DistributionResult; +import org.apache.beam.sdk.runners.TransformHierarchy; import org.apache.beam.sdk.testing.TestPipeline; import org.apache.beam.sdk.transforms.Create; import org.apache.beam.sdk.transforms.View; @@ -101,6 +105,28 @@ public class SortedBucketTransformTest { .getAll(new TupleTag("rhs")) .forEach(rhs -> outputConsumer.accept(lhs + "-" + rhs))); + private static Set transformNames(Pipeline pipeline) { + Set names = new HashSet<>(); + pipeline.traverseTopologically( + new Pipeline.PipelineVisitor.Defaults() { + private void addName(TransformHierarchy.Node node) { + names.addAll(Arrays.asList(node.getFullName().split("/"))); + } + + @Override + public CompositeBehavior enterCompositeTransform(TransformHierarchy.Node node) { + addName(node); + return CompositeBehavior.ENTER_TRANSFORM; + } + + @Override + public void visitPrimitiveTransform(TransformHierarchy.Node node) { + addName(node); + } + }); + return names; + } + private static List> makeSources(SortedBucketSource.Keying keying) { return ImmutableList.of( BucketedInput.of( @@ -274,6 +300,7 @@ private void testWithSides( new TestFileOperations(), ".txt", SortedBucketIO.DEFAULT_FILENAME_PREFIX)); + Assert.assertTrue(transformNames(transformPipeline).contains("RedistributeBuckets")); runAndValidate(targetParallelism, expectedNumBuckets, expectedWithSides); } @@ -295,6 +322,7 @@ private void test( new TestFileOperations(), ".txt", SortedBucketIO.DEFAULT_FILENAME_PREFIX)); + Assert.assertFalse(transformNames(transformPipeline).contains("RedistributeBuckets")); runAndValidate(targetParallelism, expectedNumBuckets, expected); }