Skip to content

Commit d94b56a

Browse files
committed
Improve sharding for bounded pcolleciton, to better control concurrent connections. this will keep elements for same destination close to each other and shard them. For single table write it's same behaviour, for dynamic destination it will improve reduce amount of connections used
1 parent 3d68e9d commit d94b56a

1 file changed

Lines changed: 46 additions & 1 deletion

File tree

  • sdks/java/io/google-cloud-platform/src/main/java/org/apache/beam/sdk/io/gcp/bigquery

sdks/java/io/google-cloud-platform/src/main/java/org/apache/beam/sdk/io/gcp/bigquery/StorageApiLoads.java

Lines changed: 46 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
import com.google.cloud.bigquery.storage.v1.AppendRowsRequest;
2424
import java.io.IOException;
2525
import java.nio.ByteBuffer;
26+
import java.nio.charset.StandardCharsets;
2627
import java.util.Map;
2728
import java.util.concurrent.ThreadLocalRandom;
2829
import java.util.function.Predicate;
@@ -38,6 +39,7 @@
3839
import org.apache.beam.sdk.transforms.ParDo;
3940
import org.apache.beam.sdk.transforms.Redistribute;
4041
import org.apache.beam.sdk.transforms.SerializableFunction;
42+
import org.apache.beam.sdk.transforms.Values;
4143
import org.apache.beam.sdk.transforms.errorhandling.BadRecord;
4244
import org.apache.beam.sdk.transforms.errorhandling.BadRecordRouter;
4345
import org.apache.beam.sdk.transforms.errorhandling.BadRecordRouter.ThrowingBadRecordRouter;
@@ -50,6 +52,7 @@
5052
import org.apache.beam.sdk.values.PCollectionList;
5153
import org.apache.beam.sdk.values.PCollectionTuple;
5254
import org.apache.beam.sdk.values.TupleTag;
55+
import org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.hash.Hashing;
5356
import org.joda.time.Duration;
5457

5558
/** This {@link PTransform} manages loads into BigQuery using the Storage API. */
@@ -379,12 +382,21 @@ public WriteResult expandUntriggered(
379382
PCollection<KV<DestinationT, StorageApiWritePayload>> successfulConvertedRows =
380383
convertMessagesResult.get(successfulConvertedRowsTag);
381384

382-
if (numShards > 0) {
385+
if (numShards > 0 && input.isBounded() == PCollection.IsBounded.UNBOUNDED) {
383386
successfulConvertedRows =
384387
successfulConvertedRows.apply(
385388
"ResdistibuteNumShards",
386389
Redistribute.<KV<DestinationT, StorageApiWritePayload>>arbitrarily()
387390
.withNumBuckets(numShards));
391+
} else if (numShards > 0 && input.isBounded() == PCollection.IsBounded.BOUNDED) {
392+
successfulConvertedRows =
393+
successfulConvertedRows
394+
.apply(
395+
"AddKeyWithSideInputs",
396+
ParDo.of(new AddShardKeyFn<>(dynamicDestinations, numShards))
397+
.withSideInputs(dynamicDestinations.getSideInputs()))
398+
.apply("RedistributeNumShards", Redistribute.byKey())
399+
.apply("Remove shard", Values.create());
388400
}
389401

390402
PCollectionTuple writeRecordsResult =
@@ -457,6 +469,39 @@ private void addErrorCollections(
457469
}
458470
}
459471

472+
private class AddShardKeyFn<DestinationT2, ElementT2>
473+
extends DoFn<
474+
KV<DestinationT2, StorageApiWritePayload>,
475+
KV<Integer, KV<DestinationT2, StorageApiWritePayload>>> {
476+
477+
private final StorageApiDynamicDestinations<ElementT2, DestinationT2> dynamicDestinations;
478+
private final int numShards;
479+
480+
public AddShardKeyFn(
481+
StorageApiDynamicDestinations<ElementT2, DestinationT2> dynamicDestinations,
482+
int numShards) {
483+
this.dynamicDestinations = dynamicDestinations;
484+
this.numShards = numShards;
485+
}
486+
487+
@ProcessElement
488+
public void processElement(
489+
ProcessContext c,
490+
@Element KV<DestinationT2, StorageApiWritePayload> element,
491+
OutputReceiver<KV<Integer, KV<DestinationT2, StorageApiWritePayload>>> outputReceiver) {
492+
dynamicDestinations.setSideInputAccessorFromProcessContext(c);
493+
494+
String tableUrn = dynamicDestinations.getTable(element.getKey()).getShortTableUrn();
495+
496+
int hash = Hashing.murmur3_32_fixed().hashString(tableUrn, StandardCharsets.UTF_8).asInt();
497+
498+
int shardKey =
499+
Math.floorMod(hash ^ ThreadLocalRandom.current().nextInt(numShards), numShards);
500+
501+
outputReceiver.output(KV.of(shardKey, element));
502+
}
503+
}
504+
460505
private static class ConvertInsertErrorToBadRecord
461506
extends DoFn<BigQueryStorageApiInsertError, BadRecord> {
462507

0 commit comments

Comments
 (0)