This is an automated email from the ASF dual-hosted git repository.
scwhittle 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 a3703d38100 [Dataflow Streaming][Multikey] Support MultiKey commits in
windmill clients (#38768)
a3703d38100 is described below
commit a3703d381007c54a0464bc2182559fe17b927e09
Author: Arun Pandian <[email protected]>
AuthorDate: Thu Jul 23 12:05:25 2026 -0700
[Dataflow Streaming][Multikey] Support MultiKey commits in windmill clients
(#38768)
* [Dataflow Streaming][Multikey] Support MultiKey commits in windmill
clients
- Add MultiKeyWorkItemCommitRequest to windmill.proto.
- Support MultiKey commits in Commit model and StreamingEngineWorkCommitter.
- Update GrpcCommitWorkStream to batch and stream MultiKey commit requests.
---
.../worker/windmill/client/WindmillStream.java | 5 +
.../worker/windmill/client/commits/Commit.java | 72 +++++-
.../worker/windmill/client/commits/Commits.java | 2 +-
.../windmill/client/commits/CompleteCommit.java | 15 --
.../commits/StreamingApplianceWorkCommitter.java | 10 +-
.../commits/StreamingEngineWorkCommitter.java | 92 ++++---
.../windmill/client/grpc/GrpcCommitWorkStream.java | 93 +++++--
.../dataflow/worker/FakeWindmillServer.java | 67 ++++-
.../StreamingApplianceWorkCommitterTest.java | 11 +-
.../commits/StreamingEngineWorkCommitterTest.java | 281 +++++++++++++++++++--
.../client/grpc/GrpcCommitWorkStreamTest.java | 85 +++++++
.../worker/windmill/src/main/proto/windmill.proto | 23 ++
12 files changed, 647 insertions(+), 109 deletions(-)
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/WindmillStream.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/WindmillStream.java
index 526b6789078..36001c15150 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/WindmillStream.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/WindmillStream.java
@@ -108,6 +108,11 @@ public interface WindmillStream {
Windmill.WorkItemCommitRequest request,
Consumer<Windmill.CommitStatus> onDone);
+ boolean commitMultiKeyWorkItem(
+ String computation,
+ Windmill.MultiKeyWorkItemCommitRequest request,
+ Consumer<Windmill.CommitStatus> onDone);
+
/** Flushes any pending work items to the wire. */
void flush();
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commit.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commit.java
index b840d22a343..bbd6cfc9432 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commit.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commit.java
@@ -17,35 +17,89 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;
-import com.google.auto.value.AutoValue;
+import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
+
import org.apache.beam.runners.dataflow.worker.streaming.ComputationState;
import org.apache.beam.runners.dataflow.worker.streaming.Work;
+import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.MultiKeyWorkItemCommitRequest;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest;
import org.apache.beam.sdk.annotations.Internal;
import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList;
+import org.checkerframework.checker.nullness.qual.Nullable;
/** Value class for a queued commit. */
@Internal
-@AutoValue
-public abstract class Commit {
+public class Commit {
+
+ private final ComputationState computationState;
+ private final ImmutableList<Work> workBatch;
+ private final @Nullable WorkItemCommitRequest singleKeyRequest;
+ private final @Nullable MultiKeyWorkItemCommitRequest multiKeyRequest;
public static Commit create(
WorkItemCommitRequest request, ComputationState computationState, Work
work) {
Preconditions.checkArgument(request.getSerializedSize() > 0);
- return new AutoValue_Commit(request, computationState, work);
+ return new Commit(computationState, ImmutableList.of(work), request, null);
+ }
+
+ public static Commit createMultiKey(
+ MultiKeyWorkItemCommitRequest multiKeyRequest,
+ ComputationState computationState,
+ ImmutableList<Work> workBatch) {
+ Preconditions.checkArgument(!workBatch.isEmpty());
+ return new Commit(computationState, workBatch, null, multiKeyRequest);
+ }
+
+ private Commit(
+ ComputationState computationState,
+ ImmutableList<Work> workBatch,
+ @Nullable WorkItemCommitRequest singleKeyRequest,
+ @Nullable MultiKeyWorkItemCommitRequest multiKeyRequest) {
+ this.computationState = computationState;
+ this.workBatch = workBatch;
+ this.singleKeyRequest = singleKeyRequest;
+ this.multiKeyRequest = multiKeyRequest;
}
public final String computationId() {
return computationState().getComputationId();
}
- public abstract WorkItemCommitRequest request();
+ public @Nullable WorkItemCommitRequest singleKeyRequest() {
+ return singleKeyRequest;
+ };
- public abstract ComputationState computationState();
+ public ComputationState computationState() {
+ return computationState;
+ }
+
+ public @Nullable MultiKeyWorkItemCommitRequest multiKeyRequest() {
+ return multiKeyRequest;
+ }
- public abstract Work work();
+ public ImmutableList<Work> workBatch() {
+ return workBatch;
+ }
+
+ public final int getSerializedByteSize() {
+ if (multiKeyRequest() != null) {
+ return checkStateNotNull(multiKeyRequest()).getSerializedSize();
+ }
+ return checkStateNotNull(singleKeyRequest()).getSerializedSize();
+ }
- public final int getSize() {
- return request().getSerializedSize();
+ @Override
+ public String toString() {
+ Work work = workBatch.get(0);
+ return "[computationId="
+ + computationId()
+ + ", shardingKey="
+ + work.getShardedKey()
+ + ", workId="
+ + work.id()
+ + ", workBatchSize="
+ + workBatch.size()
+ + "]";
}
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commits.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commits.java
index 498e90f78e2..0607baebeb1 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commits.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/Commits.java
@@ -31,6 +31,6 @@ public final class Commits {
private Commits() {}
public static WeightedSemaphore<Commit> maxCommitByteSemaphore() {
- return WeightedSemaphore.create(MAX_QUEUED_COMMITS_BYTES, Commit::getSize);
+ return WeightedSemaphore.create(MAX_QUEUED_COMMITS_BYTES,
Commit::getSerializedByteSize);
}
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java
index e33e853d3d7..6c0a5a98e2a 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/CompleteCommit.java
@@ -37,26 +37,11 @@ import
org.apache.beam.vendor.grpc.v1p69p0.io.grpc.stub.StreamObserver;
@AutoValue
public abstract class CompleteCommit {
- public static CompleteCommit create(Commit commit, CommitStatus
commitStatus) {
- return new AutoValue_CompleteCommit(
- commit.computationId(),
- ShardedKey.create(commit.request().getKey(),
commit.request().getShardingKey()),
- WorkId.builder()
- .setWorkToken(commit.request().getWorkToken())
- .setCacheToken(commit.request().getCacheToken())
- .build(),
- commitStatus);
- }
-
public static CompleteCommit create(
String computationId, ShardedKey shardedKey, WorkId workId, CommitStatus
status) {
return new AutoValue_CompleteCommit(computationId, shardedKey, workId,
status);
}
- public static CompleteCommit forFailedWork(Commit commit) {
- return create(commit, CommitStatus.ABORTED);
- }
-
public abstract String computationId();
public abstract ShardedKey shardedKey();
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java
index 20b95b0661d..d627490dfe9 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitter.java
@@ -17,6 +17,9 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;
+import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
+import static
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkState;
+
import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.ExecutorService;
@@ -112,7 +115,8 @@ public final class StreamingApplianceWorkCommitter
implements WorkCommitter {
}
while (commit != null) {
ComputationState computationState = commit.computationState();
- commit.work().setState(Work.State.COMMITTING);
+ checkState(commit.workBatch().size() == 1);
+ commit.workBatch().get(0).setState(Work.State.COMMITTING);
Windmill.ComputationCommitWorkRequest.Builder
computationRequestBuilder =
computationRequestMap.get(computationState);
if (computationRequestBuilder == null) {
@@ -120,10 +124,10 @@ public final class StreamingApplianceWorkCommitter
implements WorkCommitter {
computationRequestBuilder.setComputationId(computationState.getComputationId());
computationRequestMap.put(computationState,
computationRequestBuilder);
}
- computationRequestBuilder.addRequests(commit.request());
+
computationRequestBuilder.addRequests(checkStateNotNull(commit.singleKeyRequest()));
// Send the request if we've exceeded the bytes or there is no more
// pending work. commitBytes is a long, so this cannot overflow.
- commitBytes += commit.getSize();
+ commitBytes += commit.getSerializedByteSize();
if (commitBytes >= TARGET_COMMIT_BUNDLE_BYTES) {
break;
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java
index b68f53121b8..83d0dfc6cda 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java
@@ -17,6 +17,8 @@
*/
package org.apache.beam.runners.dataflow.worker.windmill.client.commits;
+import static org.apache.beam.sdk.util.Preconditions.checkStateNotNull;
+
import com.google.auto.value.AutoBuilder;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
@@ -30,6 +32,7 @@ import javax.annotation.concurrent.ThreadSafe;
import org.apache.beam.runners.dataflow.worker.streaming.WeightedBoundedQueue;
import org.apache.beam.runners.dataflow.worker.streaming.WeightedSemaphore;
import org.apache.beam.runners.dataflow.worker.streaming.Work;
+import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
import org.apache.beam.runners.dataflow.worker.windmill.client.CloseableStream;
import
org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream.CommitWorkStream;
import org.apache.beam.sdk.annotations.Internal;
@@ -100,8 +103,8 @@ public final class StreamingEngineWorkCommitter implements
WorkCommitter {
@Override
public void commit(Commit commit) {
- if (commit.work().isFailed()) {
- failCommit(commit);
+ if (shouldFailCommit(commit)) {
+ failQueuedCommit(commit);
} else {
commitQueue.put(commit);
}
@@ -109,12 +112,7 @@ public final class StreamingEngineWorkCommitter implements
WorkCommitter {
// Do this check after adding to commitQueue, else commitQueue.put() can
race with
// drainCommitQueue() in stop() and leave commits orphaned in the queue.
if (!this.isRunning.get()) {
- LOG.debug(
- "Trying to queue commit on shutdown, failing
commit=[computationId={}, shardingKey={},"
- + " workId={} ].",
- commit.computationId(),
- commit.work().getShardedKey(),
- commit.work().id());
+ LOG.debug("Trying to queue commit on shutdown, failing commit={}",
commit);
drainCommitQueue();
}
}
@@ -141,14 +139,18 @@ public final class StreamingEngineWorkCommitter
implements WorkCommitter {
private void drainCommitQueue() {
Commit queuedCommit = commitQueue.poll();
while (queuedCommit != null) {
- failCommit(queuedCommit);
+ failQueuedCommit(queuedCommit);
queuedCommit = commitQueue.poll();
}
}
- private void failCommit(Commit commit) {
- commit.work().setFailed();
- onCommitComplete.accept(CompleteCommit.forFailedWork(commit));
+ private void failQueuedCommit(Commit commit) {
+ for (Work w : commit.workBatch()) {
+ w.setFailed();
+ onCommitComplete.accept(
+ CompleteCommit.create(
+ commit.computationId(), w.getShardedKey(), w.id(),
CommitStatus.ABORTED));
+ }
}
@Override
@@ -173,8 +175,8 @@ public final class StreamingEngineWorkCommitter implements
WorkCommitter {
// take() blocks until a value is available in the commitQueue.
Preconditions.checkNotNull(initialCommit);
- if (initialCommit.work().isFailed()) {
- onCommitComplete.accept(CompleteCommit.forFailedWork(initialCommit));
+ if (shouldFailCommit(initialCommit)) {
+ failQueuedCommit(initialCommit);
initialCommit = null;
continue;
}
@@ -194,29 +196,61 @@ public final class StreamingEngineWorkCommitter
implements WorkCommitter {
}
} finally {
if (initialCommit != null) {
- failCommit(initialCommit);
+ failQueuedCommit(initialCommit);
+ }
+ }
+ }
+
+ boolean shouldFailCommit(Commit commit) {
+ for (Work w : commit.workBatch()) {
+ if (w.isFailed()) {
+ return true;
}
}
+ return false;
}
/** Adds the commit to the batch if it fits, returning true if it is
consumed. */
private boolean tryAddToCommitBatch(Commit commit,
CommitWorkStream.RequestBatcher batcher) {
Preconditions.checkNotNull(commit);
- commit.work().setState(Work.State.COMMITTING);
- activeCommitBytes.addAndGet(commit.getSize());
- boolean isCommitAccepted =
- batcher.commitWorkItem(
- commit.computationId(),
- commit.request(),
- commitStatus -> {
- onCommitComplete.accept(CompleteCommit.create(commit,
commitStatus));
- activeCommitBytes.addAndGet(-commit.getSize());
- });
+ for (Work w : commit.workBatch()) {
+ w.setState(Work.State.COMMITTING);
+ }
+ activeCommitBytes.addAndGet(commit.getSerializedByteSize());
+ boolean isCommitAccepted;
+ if (commit.multiKeyRequest() != null) {
+ isCommitAccepted =
+ batcher.commitMultiKeyWorkItem(
+ commit.computationId(),
+ checkStateNotNull(commit.multiKeyRequest()),
+ commitStatus -> {
+ for (Work w : commit.workBatch()) {
+ onCommitComplete.accept(
+ CompleteCommit.create(
+ commit.computationId(), w.getShardedKey(), w.id(),
commitStatus));
+ }
+ activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
+ });
+ } else {
+ isCommitAccepted =
+ batcher.commitWorkItem(
+ commit.computationId(),
+ checkStateNotNull(commit.singleKeyRequest()),
+ commitStatus -> {
+ Work w = commit.workBatch().get(0);
+ onCommitComplete.accept(
+ CompleteCommit.create(
+ commit.computationId(), w.getShardedKey(), w.id(),
commitStatus));
+ activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
+ });
+ }
// Since the commit was not accepted, revert the changes made above.
if (!isCommitAccepted) {
- commit.work().setState(Work.State.COMMIT_QUEUED);
- activeCommitBytes.addAndGet(-commit.getSize());
+ for (Work w : commit.workBatch()) {
+ w.setState(Work.State.COMMIT_QUEUED);
+ }
+ activeCommitBytes.addAndGet(-commit.getSerializedByteSize());
}
return isCommitAccepted;
@@ -246,8 +280,8 @@ public final class StreamingEngineWorkCommitter implements
WorkCommitter {
}
// Drop commits for failed work. Such commits will be dropped by
Windmill anyway.
- if (commit.work().isFailed()) {
- onCommitComplete.accept(CompleteCommit.forFailedWork(commit));
+ if (shouldFailCommit(commit)) {
+ failQueuedCommit(commit);
continue;
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStream.java
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStream.java
index 160b0cce013..e2a54b43cf0 100644
---
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStream.java
+++
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStream.java
@@ -34,6 +34,7 @@ import java.util.function.Consumer;
import java.util.function.Function;
import javax.annotation.Nullable;
import org.apache.beam.repackaged.core.org.apache.commons.lang3.tuple.Pair;
+import org.apache.beam.runners.dataflow.worker.windmill.Windmill;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.JobHeader;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.StreamingCommitRequestChunk;
@@ -308,7 +309,7 @@ final class GrpcCommitWorkStream
if (requests.size() == 1) {
Map.Entry<Long, PendingRequest> elem =
requests.entrySet().iterator().next();
- if (elem.getValue().getRequest().getSerializedSize()
+ if (elem.getValue().serializedCommit().size()
> AbstractWindmillStream.RPC_STREAM_CHUNK_SIZE) {
issueMultiChunkRequest(elem.getKey(), elem.getValue());
} else {
@@ -324,9 +325,10 @@ final class GrpcCommitWorkStream
StreamingCommitWorkRequest.Builder requestBuilder =
StreamingCommitWorkRequest.newBuilder();
requestBuilder
.addCommitChunkBuilder()
- .setComputationId(pendingRequest.getComputationId())
+ .setComputationId(pendingRequest.computationId())
.setRequestId(id)
.setShardingKey(pendingRequest.shardingKey())
+ .setCommitType(pendingRequest.commitType())
.setSerializedWorkItemCommit(pendingRequest.serializedCommit());
StreamingCommitWorkRequest chunk = requestBuilder.build();
synchronized (this) {
@@ -349,14 +351,15 @@ final class GrpcCommitWorkStream
for (Map.Entry<Long, PendingRequest> entry : requests.entrySet()) {
PendingRequest request = entry.getValue();
StreamingCommitRequestChunk.Builder chunkBuilder =
requestBuilder.addCommitChunkBuilder();
- if (lastComputation == null ||
!lastComputation.equals(request.getComputationId())) {
- chunkBuilder.setComputationId(request.getComputationId());
- lastComputation = request.getComputationId();
+ if (lastComputation == null ||
!lastComputation.equals(request.computationId())) {
+ chunkBuilder.setComputationId(request.computationId());
+ lastComputation = request.computationId();
}
chunkBuilder
.setRequestId(entry.getKey())
.setShardingKey(request.shardingKey())
- .setSerializedWorkItemCommit(request.serializedCommit());
+ .setSerializedWorkItemCommit(request.serializedCommit())
+ .setCommitType(request.commitType());
}
StreamingCommitWorkRequest request = requestBuilder.build();
synchronized (this) {
@@ -376,7 +379,7 @@ final class GrpcCommitWorkStream
private void issueMultiChunkRequest(long id, PendingRequest pendingRequest)
throws WindmillStreamShutdownException {
- checkNotNull(pendingRequest.getComputationId(), "Cannot commit WorkItem
w/o a computationId.");
+ checkNotNull(pendingRequest.computationId(), "Cannot commit WorkItem w/o a
computationId.");
ByteString serializedCommit = pendingRequest.serializedCommit();
synchronized (this) {
if (isShutdown) {
@@ -397,8 +400,9 @@ final class GrpcCommitWorkStream
StreamingCommitRequestChunk.newBuilder()
.setRequestId(id)
.setSerializedWorkItemCommit(chunk)
- .setComputationId(pendingRequest.getComputationId())
- .setShardingKey(pendingRequest.shardingKey());
+ .setComputationId(pendingRequest.computationId())
+ .setShardingKey(pendingRequest.shardingKey())
+ .setCommitType(pendingRequest.commitType());
int remaining = serializedCommit.size() - end;
if (remaining > 0) {
chunkBuilder.setRemainingBytesForWorkItem(remaining);
@@ -416,24 +420,44 @@ final class GrpcCommitWorkStream
private static class PendingRequest {
private final String computationId;
- private final WorkItemCommitRequest request;
+ private final long shardingKey;
+ private final ByteString serializedCommit;
+ private final StreamingCommitRequestChunk.CommitType commitType;
private final Consumer<CommitStatus> onDone;
private final long startTimeNanos; // System.nanoTime() of when request
began.
private PendingRequest(
- String computationId, WorkItemCommitRequest request,
Consumer<CommitStatus> onDone) {
+ String computationId,
+ long shardingKey,
+ ByteString serializedCommit,
+ StreamingCommitRequestChunk.CommitType commitType,
+ Consumer<CommitStatus> onDone) {
this.computationId = computationId;
- this.request = request;
+ this.shardingKey = shardingKey;
+ this.serializedCommit = serializedCommit;
+ this.commitType = commitType;
this.onDone = onDone;
this.startTimeNanos = System.nanoTime();
}
- String getComputationId() {
+ String computationId() {
return computationId;
}
- WorkItemCommitRequest getRequest() {
- return request;
+ long shardingKey() {
+ return shardingKey;
+ }
+
+ ByteString serializedCommit() {
+ return serializedCommit;
+ }
+
+ StreamingCommitRequestChunk.CommitType commitType() {
+ return commitType;
+ }
+
+ Consumer<CommitStatus> onDone() {
+ return onDone;
}
long getStartTimeNanos() {
@@ -441,21 +465,13 @@ final class GrpcCommitWorkStream
}
private long getBytes() {
- return (long) request.getSerializedSize() + computationId.length();
- }
-
- private ByteString serializedCommit() {
- return request.toByteString();
+ return (long) serializedCommit.size() + computationId.length();
}
private void completeWithStatus(CommitStatus commitStatus) {
onDone.accept(commitStatus);
}
- private long shardingKey() {
- return request.getShardingKey();
- }
-
private void abort() {
completeWithStatus(CommitStatus.ABORTED);
}
@@ -512,7 +528,34 @@ final class GrpcCommitWorkStream
return false;
}
- PendingRequest request = new PendingRequest(computation, commitRequest,
onDone);
+ PendingRequest request =
+ new PendingRequest(
+ computation,
+ commitRequest.getShardingKey(),
+ commitRequest.toByteString(),
+ StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_SINGLE_KEY,
+ onDone);
+ add(idGenerator.incrementAndGet(), request);
+ return true;
+ }
+
+ @Override
+ public boolean commitMultiKeyWorkItem(
+ String computation,
+ Windmill.MultiKeyWorkItemCommitRequest commitRequest,
+ Consumer<CommitStatus> onDone) {
+ Preconditions.checkArgument(commitRequest.getRequestsCount() > 0);
+ if (!canAccept(commitRequest.getSerializedSize() +
computation.length())) {
+ return false;
+ }
+ PendingRequest request =
+ new PendingRequest(
+ computation,
+ // Any key in the batch for routing
+ commitRequest.getRequests(0).getShardingKey(),
+ commitRequest.toByteString(),
+ StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY,
+ onDone);
add(idGenerator.incrementAndGet(), request);
return true;
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java
index 5be8ec0a6c7..e5d68376a7d 100644
---
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java
+++
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/FakeWindmillServer.java
@@ -29,7 +29,6 @@ import static org.junit.Assert.assertFalse;
import java.util.ArrayList;
import java.util.Collection;
-import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
@@ -37,10 +36,12 @@ import java.util.Queue;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentLinkedQueue;
+import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicReference;
import java.util.function.Consumer;
import java.util.function.Function;
import javax.annotation.concurrent.GuardedBy;
@@ -49,6 +50,7 @@ import
org.apache.beam.runners.dataflow.worker.streaming.WorkHeartbeatResponsePr
import org.apache.beam.runners.dataflow.worker.streaming.WorkId;
import
org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillMetadataServiceV1Alpha1Grpc;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill;
+import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitWorkResponse;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationCommitWorkRequest;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.ComputationGetDataRequest;
@@ -87,8 +89,11 @@ public final class FakeWindmillServer extends
WindmillServerStub {
private final ResponseQueue<GetDataRequest, GetDataResponse> dataToOffer;
private final ResponseQueue<Windmill.CommitWorkRequest, CommitWorkResponse>
commitsToOffer;
private final Map<WorkId, Windmill.CommitStatus> streamingCommitsToOffer;
+ private final AtomicReference<Windmill.CommitStatus>
multiKeyCommitStatusToOffer;
// Keys are work tokens.
private final Map<Long, WorkItemCommitRequest> commitsReceived;
+ private final List<Windmill.MultiKeyWorkItemCommitRequest>
multiKeyCommitsReceived =
+ new CopyOnWriteArrayList<>();
private final ArrayList<Windmill.ReportStatsRequest> statsReceived;
private final LinkedBlockingQueue<Windmill.Exception> exceptions;
private final AtomicInteger expectedExceptionCount;
@@ -118,7 +123,9 @@ public final class FakeWindmillServer extends
WindmillServerStub {
commitsToOffer =
new ResponseQueue<Windmill.CommitWorkRequest, CommitWorkResponse>()
.returnByDefault(CommitWorkResponse.getDefaultInstance());
- streamingCommitsToOffer = new HashMap<>();
+ streamingCommitsToOffer = new ConcurrentHashMap<>();
+ // Respond multikey commits with ok, unless overridden.
+ multiKeyCommitStatusToOffer = new AtomicReference<>(CommitStatus.OK);
commitsReceived = new ConcurrentHashMap<>();
exceptions = new LinkedBlockingQueue<>();
expectedExceptionCount = new AtomicInteger();
@@ -153,6 +160,11 @@ public final class FakeWindmillServer extends
WindmillServerStub {
return streamingCommitsToOffer;
}
+ /** @param commitStatus status to return to multiKeyCommits */
+ public void setMultiKeyCommitStatus(CommitStatus commitStatus) {
+ this.multiKeyCommitStatusToOffer.set(commitStatus);
+ }
+
@Override
public Windmill.GetWorkResponse getWork(Windmill.GetWorkRequest request) {
LOG.debug("getWorkRequest: {}", request.toString());
@@ -400,6 +412,7 @@ public final class FakeWindmillServer extends
WindmillServerStub {
public RequestBatcher batcher() {
return new RequestBatcher() {
final List<RequestAndDone> requests = new ArrayList<>();
+ final List<MultiKeyRequestAndDone> multiKeyRequests = new
ArrayList<>();
@Override
public boolean commitWorkItem(
@@ -423,6 +436,17 @@ public final class FakeWindmillServer extends
WindmillServerStub {
return true;
}
+ @Override
+ public boolean commitMultiKeyWorkItem(
+ String computation,
+ Windmill.MultiKeyWorkItemCommitRequest request,
+ Consumer<Windmill.CommitStatus> onDone) {
+ LOG.debug("commitWorkStream::commitMultiKeyWorkItem: {}", request);
+ multiKeyRequests.add(new MultiKeyRequestAndDone(request, onDone));
+ flush();
+ return true;
+ }
+
@Override
public void flush() {
for (RequestAndDone elem : requests) {
@@ -445,6 +469,24 @@ public final class FakeWindmillServer extends
WindmillServerStub {
.orElse(Windmill.CommitStatus.OK));
}
requests.clear();
+
+ for (MultiKeyRequestAndDone elem : multiKeyRequests) {
+ if (dropStreamingCommits) {
+ for (WorkItemCommitRequest workRequest :
elem.request.getRequestsList()) {
+ droppedStreamingCommits.put(workRequest.getWorkToken(),
elem.onDone);
+ }
+ continue;
+ }
+
+ multiKeyCommitsReceived.add(elem.request);
+ for (WorkItemCommitRequest workRequest :
elem.request.getRequestsList()) {
+ commitsReceived.put(workRequest.getWorkToken(), workRequest);
+ }
+
+ Windmill.CommitStatus status = multiKeyCommitStatusToOffer.get();
+ elem.onDone.accept(status);
+ }
+ multiKeyRequests.clear();
}
class RequestAndDone {
@@ -456,6 +498,18 @@ public final class FakeWindmillServer extends
WindmillServerStub {
this.onDone = onDone;
}
}
+
+ class MultiKeyRequestAndDone {
+ final Consumer<Windmill.CommitStatus> onDone;
+ final Windmill.MultiKeyWorkItemCommitRequest request;
+
+ MultiKeyRequestAndDone(
+ Windmill.MultiKeyWorkItemCommitRequest request,
+ Consumer<Windmill.CommitStatus> onDone) {
+ this.request = request;
+ this.onDone = onDone;
+ }
+ }
};
}
@@ -518,6 +572,15 @@ public final class FakeWindmillServer extends
WindmillServerStub {
public void clearCommitsReceived() {
commitsRequested = 0;
commitsReceived.clear();
+ multiKeyCommitsReceived.clear();
+ }
+
+ public List<Windmill.MultiKeyWorkItemCommitRequest>
getMultiKeyCommitsReceived() {
+ return multiKeyCommitsReceived;
+ }
+
+ public void clearMultiKeyCommitsReceived() {
+ multiKeyCommitsReceived.clear();
}
public ConcurrentHashMap<Long, Consumer<Windmill.CommitStatus>>
waitForDroppedCommits(
diff --git
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java
index 5c3132ae471..0596210a027 100644
---
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java
+++
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingApplianceWorkCommitterTest.java
@@ -128,10 +128,11 @@ public class StreamingApplianceWorkCommitterTest {
fakeWindmillServer.waitForAndGetCommits(commits.size());
for (Commit commit : commits) {
+ assertThat(commit.workBatch()).hasSize(1);
Windmill.WorkItemCommitRequest request =
- committed.get(commit.work().getWorkItem().getWorkToken());
+
committed.get(commit.workBatch().get(0).getWorkItem().getWorkToken());
assertNotNull(request);
- assertThat(request).isEqualTo(commit.request());
+ assertThat(request).isEqualTo(commit.singleKeyRequest());
}
assertThat(completeCommits).hasSize(commits.size());
@@ -141,12 +142,14 @@ public class StreamingApplianceWorkCommitterTest {
(CompleteCommit completeCommit, Commit commit) ->
completeCommit.computationId().equals(commit.computationId())
&& completeCommit.status() == Windmill.CommitStatus.OK
- && completeCommit.workId().equals(commit.work().id())
+ && commit.workBatch().size() == 1
+ &&
completeCommit.workId().equals(commit.workBatch().get(0).id())
&& completeCommit
.shardedKey()
.equals(
ShardedKey.create(
- commit.request().getKey(),
commit.request().getShardingKey())),
+ commit.singleKeyRequest().getKey(),
+
commit.singleKeyRequest().getShardingKey())),
"expected to equal"))
.containsExactlyElementsIn(commits);
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java
index 01197622c24..881cf620e8d 100644
---
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java
+++
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java
@@ -19,6 +19,7 @@ package
org.apache.beam.runners.dataflow.worker.windmill.client.commits;
import static com.google.common.truth.Truth.assertThat;
import static
org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus.OK;
+import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertNotNull;
import static org.junit.Assert.assertTrue;
import static org.mockito.Mockito.mock;
@@ -53,6 +54,7 @@ import org.apache.beam.runners.dataflow.worker.streaming.Work;
import org.apache.beam.runners.dataflow.worker.streaming.WorkId;
import org.apache.beam.runners.dataflow.worker.util.BoundedQueueExecutor;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill;
+import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItem;
import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.WorkItemCommitRequest;
import org.apache.beam.runners.dataflow.worker.windmill.client.CloseableStream;
@@ -62,6 +64,7 @@ import
org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDa
import
org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender;
import org.apache.beam.vendor.grpc.v1p69p0.com.google.protobuf.ByteString;
import org.apache.beam.vendor.grpc.v1p69p0.io.grpc.testing.GrpcCleanupRule;
+import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableList;
import
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.ImmutableMap;
import org.joda.time.Duration;
import org.joda.time.Instant;
@@ -134,12 +137,10 @@ public class StreamingEngineWorkCommitterTest {
null);
}
- private static CompleteCommit asCompleteCommit(Commit commit,
Windmill.CommitStatus status) {
- if (commit.work().isFailed()) {
- return CompleteCommit.forFailedWork(commit);
- }
-
- return CompleteCommit.create(commit, status);
+ private static CompleteCommit asCompleteCommit(
+ String computationId, Work work, Windmill.CommitStatus status) {
+ Windmill.CommitStatus finalStatus = work.isFailed() ?
Windmill.CommitStatus.ABORTED : status;
+ return CompleteCommit.create(computationId, work.getShardedKey(),
work.id(), finalStatus);
}
@Before
@@ -186,10 +187,15 @@ public class StreamingEngineWorkCommitterTest {
waitForExpectedSetSize(completeCommits, 5);
for (Commit commit : commits) {
- WorkItemCommitRequest request =
committed.get(commit.work().getWorkItem().getWorkToken());
+ assertThat(commit.workBatch()).hasSize(1);
+ WorkItemCommitRequest request =
+
committed.get(commit.workBatch().get(0).getWorkItem().getWorkToken());
assertNotNull(request);
- assertThat(request).isEqualTo(commit.request());
- assertThat(completeCommits).contains(asCompleteCommit(commit,
Windmill.CommitStatus.OK));
+ assertThat(request).isEqualTo(commit.singleKeyRequest());
+ assertThat(completeCommits)
+ .contains(
+ asCompleteCommit(
+ commit.computationId(), commit.workBatch().get(0),
Windmill.CommitStatus.OK));
}
workCommitter.stop();
@@ -224,14 +230,24 @@ public class StreamingEngineWorkCommitterTest {
waitForExpectedSetSize(completeCommits, 10);
for (Commit commit : commits) {
- if (commit.work().isFailed()) {
+ assertThat(commit.workBatch()).hasSize(1);
+ if (commit.workBatch().get(0).isFailed()) {
assertThat(completeCommits)
- .contains(asCompleteCommit(commit, Windmill.CommitStatus.ABORTED));
-
assertThat(committed).doesNotContainKey(commit.work().getWorkItem().getWorkToken());
+ .contains(
+ asCompleteCommit(
+ commit.computationId(),
+ commit.workBatch().get(0),
+ Windmill.CommitStatus.ABORTED));
+ assertThat(committed)
+
.doesNotContainKey(commit.workBatch().get(0).getWorkItem().getWorkToken());
} else {
- assertThat(completeCommits).contains(asCompleteCommit(commit,
Windmill.CommitStatus.OK));
+ assertThat(completeCommits)
+ .contains(
+ asCompleteCommit(
+ commit.computationId(), commit.workBatch().get(0),
Windmill.CommitStatus.OK));
assertThat(committed)
- .containsEntry(commit.work().getWorkItem().getWorkToken(),
commit.request());
+ .containsEntry(
+ commit.workBatch().get(0).getWorkItem().getWorkToken(),
commit.singleKeyRequest());
}
}
@@ -282,11 +298,17 @@ public class StreamingEngineWorkCommitterTest {
waitForExpectedSetSize(completeCommits, commits.size());
for (Commit commit : commits) {
- WorkItemCommitRequest request =
committed.get(commit.work().getWorkItem().getWorkToken());
+ assertThat(commit.workBatch()).hasSize(1);
+ WorkItemCommitRequest request =
+
committed.get(commit.workBatch().get(0).getWorkItem().getWorkToken());
assertNotNull(request);
- assertThat(request).isEqualTo(commit.request());
+ assertThat(request).isEqualTo(commit.singleKeyRequest());
assertThat(completeCommits)
- .contains(asCompleteCommit(commit,
expectedCommitStatus.get(commit.work().id())));
+ .contains(
+ asCompleteCommit(
+ commit.computationId(),
+ commit.workBatch().get(0),
+ expectedCommitStatus.get(commit.workBatch().get(0).id())));
}
workCommitter.stop();
@@ -313,6 +335,14 @@ public class StreamingEngineWorkCommitterTest {
return false;
}
+ @Override
+ public boolean commitMultiKeyWorkItem(
+ String computation,
+ Windmill.MultiKeyWorkItemCommitRequest request,
+ Consumer<Windmill.CommitStatus> onDone) {
+ return false;
+ }
+
@Override
public void flush() {}
};
@@ -370,7 +400,8 @@ public class StreamingEngineWorkCommitterTest {
}
for (Commit commit : commits) {
- assertTrue(commit.work().isFailed());
+ assertThat(commit.workBatch()).hasSize(1);
+ assertTrue(commit.workBatch().get(0).isFailed());
}
}
@@ -409,10 +440,15 @@ public class StreamingEngineWorkCommitterTest {
waitForExpectedSetSize(completeCommits, commits.size());
for (Commit commit : commits) {
- WorkItemCommitRequest request =
committed.get(commit.work().getWorkItem().getWorkToken());
+ assertThat(commit.workBatch()).hasSize(1);
+ WorkItemCommitRequest request =
+
committed.get(commit.workBatch().get(0).getWorkItem().getWorkToken());
assertNotNull(request);
- assertThat(request).isEqualTo(commit.request());
- assertThat(completeCommits).contains(asCompleteCommit(commit,
Windmill.CommitStatus.OK));
+ assertThat(request).isEqualTo(commit.singleKeyRequest());
+ assertThat(completeCommits)
+ .contains(
+ asCompleteCommit(
+ commit.computationId(), commit.workBatch().get(0),
Windmill.CommitStatus.OK));
}
workCommitter.stop();
@@ -474,4 +510,207 @@ public class StreamingEngineWorkCommitterTest {
waitForExpectedSetSize(completeCommits, sentCommits.intValue());
}
+
+ @Test
+ public void testCommit_multiKeyCommitSuccess() {
+ Set<CompleteCommit> completeCommits = Collections.newSetFromMap(new
ConcurrentHashMap<>());
+ workCommitter = createWorkCommitter(completeCommits::add);
+
+ Work workA = createMockWork(101L);
+ Work workB = createMockWork(102L);
+ Work workC = createMockWork(103L);
+
+ Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest =
+ Windmill.MultiKeyWorkItemCommitRequest.newBuilder()
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workA.getWorkItem().getKey())
+ .setShardingKey(workA.getWorkItem().getShardingKey())
+ .setWorkToken(workA.getWorkItem().getWorkToken())
+ .setCacheToken(workA.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workB.getWorkItem().getKey())
+ .setShardingKey(workB.getWorkItem().getShardingKey())
+ .setWorkToken(workB.getWorkItem().getWorkToken())
+ .setCacheToken(workB.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workC.getWorkItem().getKey())
+ .setShardingKey(workC.getWorkItem().getShardingKey())
+ .setWorkToken(workC.getWorkItem().getWorkToken())
+ .setCacheToken(workC.getWorkItem().getCacheToken())
+ .build())
+ .build();
+
+ Commit commit =
+ Commit.createMultiKey(
+ multiKeyRequest,
+ createComputationState("computationId"),
+ ImmutableList.of(workA, workB, workC));
+
+ workCommitter.start();
+ workCommitter.commit(commit);
+
+ // Wait for the server to receive and process the commits
+ fakeWindmillServer.waitForAndGetCommits(3);
+ waitForExpectedSetSize(completeCommits, 3);
+
+ // Verify that FakeWindmillServer received all 3 work requests in
multiKeyCommitsReceived
+ List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ fakeWindmillServer.getMultiKeyCommitsReceived();
+ assertThat(multiKeyCommits).hasSize(1);
+ assertThat(multiKeyCommits.get(0)).isEqualTo(multiKeyRequest);
+
+ // Verify all three works are completed successfully
+ assertThat(completeCommits)
+ .containsExactly(
+ CompleteCommit.create(
+ "computationId", workA.getShardedKey(), workA.id(),
CommitStatus.OK),
+ CompleteCommit.create(
+ "computationId", workB.getShardedKey(), workB.id(),
CommitStatus.OK),
+ CompleteCommit.create(
+ "computationId", workC.getShardedKey(), workC.id(),
CommitStatus.OK));
+
+ // There should be no more commits in the queue
+ assertEquals(0, workCommitter.currentActiveCommitBytes());
+ workCommitter.stop();
+ }
+
+ @Test
+ public void testCommit_multiKeyCommitFailedWork() {
+ Set<CompleteCommit> completeCommits = Collections.newSetFromMap(new
ConcurrentHashMap<>());
+ workCommitter = createWorkCommitter(completeCommits::add);
+
+ Work workA = createMockWork(101L);
+ Work workB = createMockWork(102L);
+ Work workC = createMockWork(103L);
+
+ // Mark non-primary key B as failed
+ workB.setFailed();
+
+ Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest =
+ Windmill.MultiKeyWorkItemCommitRequest.newBuilder()
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workA.getWorkItem().getKey())
+ .setShardingKey(workA.getWorkItem().getShardingKey())
+ .setWorkToken(workA.getWorkItem().getWorkToken())
+ .setCacheToken(workA.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workB.getWorkItem().getKey())
+ .setShardingKey(workB.getWorkItem().getShardingKey())
+ .setWorkToken(workB.getWorkItem().getWorkToken())
+ .setCacheToken(workB.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workC.getWorkItem().getKey())
+ .setShardingKey(workC.getWorkItem().getShardingKey())
+ .setWorkToken(workC.getWorkItem().getWorkToken())
+ .setCacheToken(workC.getWorkItem().getCacheToken())
+ .build())
+ .build();
+
+ Commit commit =
+ Commit.createMultiKey(
+ multiKeyRequest,
+ createComputationState("computationId"),
+ ImmutableList.of(workA, workB, workC));
+
+ workCommitter.start();
+ workCommitter.commit(commit);
+
+ // The entire batch must be aborted immediately without making network
calls
+ waitForExpectedSetSize(completeCommits, 3);
+
+ // Verify all three works are aborted individually
+ assertThat(completeCommits)
+ .containsExactly(
+ CompleteCommit.create(
+ "computationId", workA.getShardedKey(), workA.id(),
CommitStatus.ABORTED),
+ CompleteCommit.create(
+ "computationId", workB.getShardedKey(), workB.id(),
CommitStatus.ABORTED),
+ CompleteCommit.create(
+ "computationId", workC.getShardedKey(), workC.id(),
CommitStatus.ABORTED));
+
+ // There should be no more commits in the queue
+ assertEquals(0, workCommitter.currentActiveCommitBytes());
+ workCommitter.stop();
+ }
+
+ @Test
+ public void testCommit_multiKeyCommitStatusNotOK() {
+ Set<CompleteCommit> completeCommits = Collections.newSetFromMap(new
ConcurrentHashMap<>());
+ workCommitter = createWorkCommitter(completeCommits::add);
+
+ Work workA = createMockWork(101L);
+ Work workB = createMockWork(102L);
+ Work workC = createMockWork(103L);
+
+ Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest =
+ Windmill.MultiKeyWorkItemCommitRequest.newBuilder()
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workA.getWorkItem().getKey())
+ .setShardingKey(workA.getWorkItem().getShardingKey())
+ .setWorkToken(workA.getWorkItem().getWorkToken())
+ .setCacheToken(workA.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workB.getWorkItem().getKey())
+ .setShardingKey(workB.getWorkItem().getShardingKey())
+ .setWorkToken(workB.getWorkItem().getWorkToken())
+ .setCacheToken(workB.getWorkItem().getCacheToken())
+ .build())
+ .addRequests(
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(workC.getWorkItem().getKey())
+ .setShardingKey(workC.getWorkItem().getShardingKey())
+ .setWorkToken(workC.getWorkItem().getWorkToken())
+ .setCacheToken(workC.getWorkItem().getCacheToken())
+ .build())
+ .build();
+
+ Commit commit =
+ Commit.createMultiKey(
+ multiKeyRequest,
+ createComputationState("computationId"),
+ ImmutableList.of(workA, workB, workC));
+
+ // Respond to multi key commit with NOT_FOUND status.
+ fakeWindmillServer.setMultiKeyCommitStatus(CommitStatus.NOT_FOUND);
+
+ workCommitter.start();
+ workCommitter.commit(commit);
+
+ // Wait for the server to receive and process the commits
+ fakeWindmillServer.waitForAndGetCommits(3);
+ waitForExpectedSetSize(completeCommits, 3);
+
+ // Verify that FakeWindmillServer received the multi-key commit
+ List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+ fakeWindmillServer.getMultiKeyCommitsReceived();
+ assertThat(multiKeyCommits).hasSize(1);
+ assertThat(multiKeyCommits.get(0)).isEqualTo(multiKeyRequest);
+
+ // Verify all three works in the multi-key commit are completed with
NOT_FOUND status
+ assertThat(completeCommits)
+ .containsExactly(
+ CompleteCommit.create(
+ "computationId", workA.getShardedKey(), workA.id(),
CommitStatus.NOT_FOUND),
+ CompleteCommit.create(
+ "computationId", workB.getShardedKey(), workB.id(),
CommitStatus.NOT_FOUND),
+ CompleteCommit.create(
+ "computationId", workC.getShardedKey(), workC.id(),
CommitStatus.NOT_FOUND));
+
+ // There should be no more commits in the queue
+ assertEquals(0, workCommitter.currentActiveCommitBytes());
+ workCommitter.stop();
+ }
}
diff --git
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java
index 9c3d5c9c3ef..1e995f4047c 100644
---
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java
+++
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcCommitWorkStreamTest.java
@@ -42,6 +42,8 @@ import java.util.concurrent.atomic.AtomicReference;
import java.util.function.Supplier;
import
org.apache.beam.runners.dataflow.worker.windmill.CloudWindmillServiceV1Alpha1Grpc;
import org.apache.beam.runners.dataflow.worker.windmill.Windmill;
+import org.apache.beam.runners.dataflow.worker.windmill.Windmill.CommitStatus;
+import
org.apache.beam.runners.dataflow.worker.windmill.Windmill.StreamingCommitResponse;
import org.apache.beam.runners.dataflow.worker.windmill.WindmillConnection;
import
org.apache.beam.runners.dataflow.worker.windmill.client.TriggeredScheduledExecutorService;
import org.apache.beam.runners.dataflow.worker.windmill.client.WindmillStream;
@@ -1134,6 +1136,89 @@ public class GrpcCommitWorkStreamTest {
assertTrue(commitWorkStream.awaitTermination(10, TimeUnit.SECONDS));
}
+ @Test
+ public void testCommit_multiKeyCommit() throws Exception {
+ testMultiKeyCommit(CommitStatus.OK);
+ }
+
+ @Test
+ public void testCommit_multiKeyCommit_Failure() throws Exception {
+ testMultiKeyCommit(CommitStatus.NOT_FOUND);
+ }
+
+ private void testMultiKeyCommit(CommitStatus commitStatus) throws Exception {
+ GrpcCommitWorkStream commitWorkStream = createCommitWorkStream();
+ FakeWindmillGrpcService.CommitStreamInfo streamInfo =
waitForConnectionAndConsumeHeader();
+
+ CompletableFuture<CommitStatus> commitStatusFuture = new
CompletableFuture<>();
+
+ // 1. Construct two individual WorkItemCommitRequests
+ long shardingKey1 = 101L;
+ long workToken1 = 201L;
+ long cacheToken1 = 301L;
+ long shardingKey2 = 102L;
+ long workToken2 = 202L;
+ long cacheToken2 = 302L;
+ Windmill.WorkItemCommitRequest request1 =
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(ByteString.copyFromUtf8("key1"))
+ .setShardingKey(shardingKey1)
+ .setWorkToken(workToken1)
+ .setCacheToken(cacheToken1)
+ .build();
+ Windmill.WorkItemCommitRequest request2 =
+ Windmill.WorkItemCommitRequest.newBuilder()
+ .setKey(ByteString.copyFromUtf8("key2"))
+ .setShardingKey(shardingKey2)
+ .setWorkToken(workToken2)
+ .setCacheToken(cacheToken2)
+ .build();
+
+ // 2. Wrap them into a MultiKeyWorkItemCommitRequest
+ Windmill.MultiKeyWorkItemCommitRequest multiKeyRequest =
+ Windmill.MultiKeyWorkItemCommitRequest.newBuilder()
+ .addRequests(request1)
+ .addRequests(request2)
+ .build();
+
+ // 3. Commit the multi-key work item using the request batcher
+ try (WindmillStream.CommitWorkStream.RequestBatcher batcher =
commitWorkStream.batcher()) {
+ assertTrue(
+ batcher.commitMultiKeyWorkItem(
+ COMPUTATION_ID, multiKeyRequest, commitStatusFuture::complete));
+ }
+
+ // 4. Receive and assert request properties on FakeWindmillGrpcService
+ Windmill.StreamingCommitWorkRequest request = streamInfo.requests.take();
+ assertThat(request.getCommitChunkCount()).isEqualTo(1);
+
+ Windmill.StreamingCommitRequestChunk chunk = request.getCommitChunk(0);
+
+ // Assert that the commit type is correctly identified as
COMMIT_TYPE_MULTI_KEY
+ assertThat(chunk.getCommitType())
+
.isEqualTo(Windmill.StreamingCommitRequestChunk.CommitType.COMMIT_TYPE_MULTI_KEY);
+
+ // Assert that the routing sharding key is mapped to the first request's
sharding key
+ assertThat(chunk.getShardingKey()).isEqualTo(request1.getShardingKey());
+
+ // Assert that the serialized payload matches the input multiKeyRequest
+ Windmill.MultiKeyWorkItemCommitRequest parsedRequest =
+
Windmill.MultiKeyWorkItemCommitRequest.parseFrom(chunk.getSerializedWorkItemCommit());
+ assertThat(parsedRequest).isEqualTo(multiKeyRequest);
+
+ // 5. Respond with the generated requestId to complete the commit
+ long requestId = chunk.getRequestId();
+ StreamingCommitResponse.Builder builder =
+ StreamingCommitResponse.newBuilder().addRequestId(requestId);
+ if (commitStatus != CommitStatus.OK) {
+ builder.addStatus(commitStatus);
+ }
+ streamInfo.responseObserver.onNext(builder.build());
+
+ // 6. Verify callback completed with expected sCommitStatus
+ assertThat(commitStatusFuture.get()).isEqualTo(commitStatus);
+ }
+
@Test
public void testCommitWorkItem_stopsRetriesAfterDuration() throws Exception {
int numCommits = 1;
diff --git
a/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto
b/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto
index aaa09c105fc..a7a99e2ca5a 100644
---
a/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto
+++
b/runners/google-cloud-dataflow-java/worker/windmill/src/main/proto/windmill.proto
@@ -678,9 +678,24 @@ message WorkItemCommitRequest {
reserved 6, 23;
}
+message MultiKeyWorkItemCommitRequest {
+ optional Uint128Proto key_group = 7;
+
+ repeated WorkItemCommitRequest requests = 1;
+
+ repeated OutputMessageBundle output_messages = 2;
+
+ repeated PubSubMessageBundle pubsub_messages = 3;
+
+ repeated int64 finalize_ids = 4 [packed = true];
+
+ reserved 6;
+}
+
message ComputationCommitWorkRequest {
required string computation_id = 1;
repeated WorkItemCommitRequest requests = 2;
+ repeated MultiKeyWorkItemCommitRequest multi_key_requests = 3;
}
message CommitWorkRequest {
@@ -906,6 +921,14 @@ message StreamingCommitRequestChunk {
// before handing off to the WindmillHost for processing.
optional int64 remaining_bytes_for_work_item = 4;
optional bytes serialized_work_item_commit = 5;
+
+ enum CommitType {
+ COMMIT_TYPE_UNSPECIFIED = 0;
+ COMMIT_TYPE_SINGLE_KEY = 1;
+ COMMIT_TYPE_MULTI_KEY = 2;
+ }
+
+ optional CommitType commit_type = 7;
}
message StreamingCommitResponse {