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 {

Reply via email to