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 92e3f353a85 [Dataflow Streaming] [Multi Key] Flush streaming sinks at 
key boundaries and bundle completion (#39961)
92e3f353a85 is described below

commit 92e3f353a85e11ddcdb9ae948511e7b8b662a762
Author: Arun Pandian <[email protected]>
AuthorDate: Mon Sep 21 01:57:08 2026 -0700

    [Dataflow Streaming] [Multi Key] Flush streaming sinks at key boundaries 
and bundle completion (#39961)
    
    When multi-key bundling is enabled in Dataflow Streaming Engine,
    outputs (productions) emitted during a key's processElement lifecycle
    must be attributed to that specific key's WorkItemCommitRequest, while
    outputs emitted during DoFn finishBundle must be attributed to the
    bundle-level MultiKeyWorkItemCommitRequest.
    
    This change:
    1. Adds finishKey(@Nullable Object key) to the SinkWriter interface
       and implements it across streaming sinks
    2. In WindmillSink, PubsubSink, and PubsubDynamicSink:
       - When multi-key bundles are enabled: flushes to the active key's commit 
builder
         in finishKey; flushes remaining finishBundle outputs to bundle-level
         outputs (bundleOutputMessages / bundlePubsubMessages) in close().
       - When multi-key bundles are disabled: finishKey does not flush; close() 
flushes
         all outputs directly into the single key's commit request.
       - Updates abort() to cleanly discard buffered messages and reset
         internal streams without flushing.
    3. In StreamingModeExecutionContext and StreamingWorkScheduler, tracks
       bundleOutputMessages and bundlePubsubMessages and attaches them to
       MultiKeyWorkItemCommitRequest.
---
 .../runners/dataflow/worker/PubsubDynamicSink.java |  45 +-
 .../beam/runners/dataflow/worker/PubsubSink.java   |  46 +-
 .../dataflow/worker/SizeReportingSinkWrapper.java  |   6 +
 .../worker/StreamingModeExecutionContext.java      | 168 +++++---
 .../beam/runners/dataflow/worker/WindmillSink.java |  27 +-
 .../dataflow/worker/util/common/worker/Sink.java   |   6 +
 .../worker/util/common/worker/WriteOperation.java  |   9 +-
 .../work/processing/ExecuteWorkResult.java         |  57 +++
 .../work/processing/StreamingWorkScheduler.java    |  72 ++--
 .../dataflow/worker/PubsubDynamicSinkTest.java     | 210 ++++++++-
 .../runners/dataflow/worker/PubsubSinkTest.java    | 174 +++++++-
 .../worker/StreamingDataflowWorkerTest.java        | 467 +++++++++++++++++++++
 .../worker/StreamingModeExecutionContextTest.java  |  15 +-
 .../dataflow/worker/WorkerCustomSourcesTest.java   |  38 +-
 .../util/common/worker/WriteOperationTest.java     |  59 +++
 15 files changed, 1235 insertions(+), 164 deletions(-)

diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSink.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSink.java
index 9098cdc6717..a9f9f24f4a5 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSink.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSink.java
@@ -113,9 +113,12 @@ public class PubsubDynamicSink extends 
Sink<WindowedValue<PubsubMessage>> {
       return stream.toByteStringAndReset();
     }
 
-    public void close(Windmill.PubSubMessageBundle.Builder outputBuilder) 
throws IOException {
-      context.getOutputBuilder().addPubsubMessages(outputBuilder);
-      outputBuilder.clear();
+    private Windmill.PubSubMessageBundle.Builder createOutputBuilder(String 
topic) {
+      return Windmill.PubSubMessageBundle.newBuilder()
+          .setTopic(topic)
+          .setTimestampLabel(timestampLabel)
+          .setIdLabel(idLabel)
+          .setWithAttributes(true);
     }
 
     @Override
@@ -127,32 +130,44 @@ public class PubsubDynamicSink extends 
Sink<WindowedValue<PubsubMessage>> {
           !dataTopic.isEmpty(), "No topic set for message when using dynamic 
topics.");
       ByteString byteString = getDataFromMessage(data.getValue(), stream);
       Windmill.PubSubMessageBundle.Builder builder =
-          outputBuilders.computeIfAbsent(
-              dataTopic,
-              topic ->
-                  context
-                      .getOutputBuilder()
-                      .addPubsubMessagesBuilder()
-                      .setTopic(topic)
-                      .setTimestampLabel(timestampLabel)
-                      .setIdLabel(idLabel)
-                      .setWithAttributes(true));
+          outputBuilders.computeIfAbsent(dataTopic, this::createOutputBuilder);
       builder.addMessages(
           Windmill.Message.newBuilder()
               .setData(byteString)
               
.setTimestamp(WindmillTimeUtils.harnessToWindmillTimestamp(data.getTimestamp()))
               .build());
+
       return byteString.size();
     }
 
+    private void flush(boolean bundleLevel) {
+      try {
+        for (Windmill.PubSubMessageBundle.Builder builder : 
outputBuilders.values()) {
+          if (builder.getMessagesCount() > 0) {
+            Windmill.PubSubMessageBundle pubsubMessages = builder.build();
+            if (bundleLevel) {
+              // If/when we add support for ordering keys, the flush needs to 
happen at the key
+              // level
+              context.addBundlePubsubMessages(pubsubMessages);
+            } else {
+              context.getKeyOutputBuilder().addPubsubMessages(pubsubMessages);
+            }
+          }
+        }
+      } finally {
+        outputBuilders.clear();
+      }
+    }
+
     @Override
     public void close() throws IOException {
-      outputBuilders.clear();
+      flush(/* bundleLevel= */ context.multiKeyBundleEnabled());
     }
 
     @Override
     public void abort() throws IOException {
-      close();
+      outputBuilders.clear();
+      stream.reset();
     }
   }
 
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubSink.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubSink.java
index 2f4b26b89ab..a3c8fa94f96 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubSink.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/PubsubSink.java
@@ -135,7 +135,7 @@ class PubsubSink<T> extends Sink<WindowedValue<T>> {
 
   @Override
   public SinkWriter<WindowedValue<T>> writer() {
-    return new PubsubWriter(topic);
+    return new PubsubWriter();
   }
 
   /** The SinkWriter for a PubsubSink. */
@@ -143,16 +143,19 @@ class PubsubSink<T> extends Sink<WindowedValue<T>> {
     private Windmill.PubSubMessageBundle.Builder outputBuilder;
     private ByteStringOutputStream stream; // Kept across adds for buffer 
reuse.
 
-    private PubsubWriter(String topic) {
-      outputBuilder =
-          Windmill.PubSubMessageBundle.newBuilder()
-              .setTopic(topic)
-              .setTimestampLabel(timestampLabel)
-              .setIdLabel(idLabel)
-              .setWithAttributes(withAttributes);
+    private PubsubWriter() {
+      outputBuilder = createOutputBuilder();
       stream = new ByteStringOutputStream();
     }
 
+    private Windmill.PubSubMessageBundle.Builder createOutputBuilder() {
+      return Windmill.PubSubMessageBundle.newBuilder()
+          .setTopic(topic)
+          .setTimestampLabel(timestampLabel)
+          .setIdLabel(idLabel)
+          .setWithAttributes(withAttributes);
+    }
+
     @Override
     public long add(WindowedValue<T> data) throws IOException {
       if (!stream.isEmpty()) {
@@ -187,18 +190,33 @@ class PubsubSink<T> extends Sink<WindowedValue<T>> {
       return byteString.size();
     }
 
+    private void flush(boolean bundleLevel) {
+      try {
+        Windmill.PubSubMessageBundle pubsubMessages = outputBuilder.build();
+        if (pubsubMessages.getMessagesCount() > 0) {
+          if (bundleLevel) {
+            // If/when we add support for ordering keys, the flush needs to 
happen at the key level
+            context.addBundlePubsubMessages(pubsubMessages);
+          } else {
+            context.getKeyOutputBuilder().addPubsubMessages(pubsubMessages);
+          }
+        }
+      } finally {
+        // TODO: Set to createOutputBuilder() for if/when adding support to 
reuse the sink across
+        // bundles.
+        outputBuilder.clear();
+      }
+    }
+
     @Override
     public void close() throws IOException {
-      Windmill.PubSubMessageBundle pubsubMessages = outputBuilder.build();
-      if (pubsubMessages.getMessagesCount() > 0) {
-        context.getOutputBuilder().addPubsubMessages(pubsubMessages);
-      }
-      outputBuilder.clear();
+      flush(/* bundleLevel= */ context.multiKeyBundleEnabled());
     }
 
     @Override
     public void abort() throws IOException {
-      close();
+      outputBuilder.clear();
+      stream.reset();
     }
   }
 
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/SizeReportingSinkWrapper.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/SizeReportingSinkWrapper.java
index ed6a77f6d9c..1a9919baf49 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/SizeReportingSinkWrapper.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/SizeReportingSinkWrapper.java
@@ -20,6 +20,7 @@ package org.apache.beam.runners.dataflow.worker;
 import java.io.IOException;
 import org.apache.beam.runners.dataflow.worker.util.common.worker.Sink;
 import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.annotations.VisibleForTesting;
+import org.checkerframework.checker.nullness.qual.Nullable;
 
 /**
  * A wrapper for Sink that reports bytes buffered (or written) to {@link 
DataflowExecutionContext}.
@@ -65,6 +66,11 @@ public class SizeReportingSinkWrapper<T> extends Sink<T> {
       return size;
     }
 
+    @Override
+    public void finishKey(@Nullable Object key) throws IOException {
+      underlyingWriter.finishKey(key);
+    }
+
     @Override
     public void close() throws IOException {
       underlyingWriter.close();
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java
index a6d1cc99d8e..c7036aed0ba 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java
@@ -82,6 +82,7 @@ import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodin
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV1;
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV2;
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTimerData;
+import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.ExecuteWorkResult;
 import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker;
 import org.apache.beam.sdk.annotations.Internal;
 import org.apache.beam.sdk.coders.Coder;
@@ -163,7 +164,7 @@ public class StreamingModeExecutionContext
   // be used for processing many work items and these values can change during 
the context's
   // lifetime. start() is called for each work item.
   private OperationalLimits operationalLimits;
-  private Windmill.WorkItemCommitRequest.@Nullable Builder outputBuilder;
+  private Windmill.WorkItemCommitRequest.@Nullable Builder keyOutputBuilder;
 
   /**
    * Current reader used for processing {@link Work}. Set by calling {@link
@@ -195,10 +196,12 @@ public class StreamingModeExecutionContext
   private @Nullable KeyTransitionListener keyTransitionListener;
   private @Nullable FailedWorkHandler onFailedWorkHandler;
 
-  private List<Windmill.WorkItemCommitRequest.Builder> outputBuilders = 
Collections.emptyList();
+  private @Nullable List<Windmill.WorkItemCommitRequest> workItemCommits = 
null;
+  private @Nullable List<Windmill.OutputMessageBundle> bundleOutputMessages = 
null;
+  private @Nullable List<Windmill.PubSubMessageBundle> bundlePubsubMessages = 
null;
 
   // Map<finalizerId, Pair<callbackExpiration, callback>>
-  private Map<Long, Pair<Instant, Runnable>> finalizationCallbacks = 
Collections.emptyMap();
+  private @Nullable Map<Long, Pair<Instant, Runnable>> finalizationCallbacks = 
null;
   private AtomicBoolean workBatchFailed = new AtomicBoolean(false);
   private @Nullable WindmillStateReader activeStateReader;
   private long stateBytesRead = 0;
@@ -320,8 +323,10 @@ public class StreamingModeExecutionContext
   public void reset() {
     // these lists and maps are returned to callers after processing
     // don't clear and reuse, instead reset the reference.
-    this.outputBuilders = Collections.emptyList();
-    this.finalizationCallbacks = Collections.emptyMap();
+    this.workItemCommits = null;
+    this.bundleOutputMessages = null;
+    this.bundlePubsubMessages = null;
+    this.finalizationCallbacks = null;
     // Work from prior bundles might have a reference to the old 
workBatchFailed.
     // If the work gets retried it'll get the new workBatchFailed to notify 
failure.
     this.workBatchFailed = new AtomicBoolean(false);
@@ -336,7 +341,7 @@ public class StreamingModeExecutionContext
     this.onFailedWorkHandler = null;
     this.work = null;
     this.key = null;
-    this.outputBuilder = null;
+    this.keyOutputBuilder = null;
     this.sideInputStateFetcher = null;
     this.backlogBytes = UnboundedReader.BACKLOG_UNKNOWN;
     clearSinkFullHint();
@@ -353,8 +358,6 @@ public class StreamingModeExecutionContext
       FailedWorkHandler onFailedWorkHandler)
       throws CoderException {
     reset();
-    this.outputBuilders = new ArrayList<>();
-    this.finalizationCallbacks = new HashMap<>();
     this.keyCoder = keyCoder;
     this.workExecutor = workExecutor;
     this.workQueueExecutor = workQueueExecutor;
@@ -441,9 +444,28 @@ public class StreamingModeExecutionContext
     }
   }
 
-  public void flushState() {
-    checkState(finishKeyCalled, "finishKey must be called before flushState");
+  public ExecuteWorkResult flushStateAndReset() {
+    checkState(finishKeyCalled, "finishKey must be called before 
flushStateAndReset");
     flushStateInternal();
+
+    List<Windmill.WorkItemCommitRequest> workItemCommits =
+        this.workItemCommits != null ? this.workItemCommits : 
Collections.emptyList();
+    List<Windmill.OutputMessageBundle> bundleOutputMessages =
+        this.bundleOutputMessages != null ? this.bundleOutputMessages : 
Collections.emptyList();
+    List<Windmill.PubSubMessageBundle> bundlePubsubMessages =
+        this.bundlePubsubMessages != null ? this.bundlePubsubMessages : 
Collections.emptyList();
+    Map<Long, Pair<Instant, Runnable>> finalizationCallbacks =
+        this.finalizationCallbacks != null ? this.finalizationCallbacks : 
Collections.emptyMap();
+    long stateBytesRead = this.stateBytesRead;
+
+    reset();
+
+    return ExecuteWorkResult.create(
+        workItemCommits,
+        bundleOutputMessages,
+        bundlePubsubMessages,
+        finalizationCallbacks,
+        stateBytesRead);
   }
 
   /**
@@ -557,8 +579,8 @@ public class StreamingModeExecutionContext
     return getWorkItem().getTimers().getTimersList();
   }
 
-  public Windmill.WorkItemCommitRequest.Builder getOutputBuilder() {
-    return checkStateNotNull(outputBuilder);
+  public Windmill.WorkItemCommitRequest.Builder getKeyOutputBuilder() {
+    return checkStateNotNull(keyOutputBuilder);
   }
 
   /**
@@ -612,47 +634,50 @@ public class StreamingModeExecutionContext
   }
 
   private void flushStateInternal() {
-    Map<Long, Pair<Instant, Runnable>> callbacks = new HashMap<>();
-
     for (StepContext stepContext : getAllStepContexts()) {
       stepContext.flushState();
-      for (Pair<Instant, BundleFinalizer.Callback> bundleFinalizer :
-          stepContext.flushBundleFinalizerCallbacks()) {
-        long id = ThreadLocalRandom.current().nextLong();
-        callbacks.put(
-            id,
-            Pair.of(
-                bundleFinalizer.getLeft(),
-                () -> {
-                  try {
-                    bundleFinalizer.getRight().onBundleSuccess();
-                  } catch (Exception e) {
-                    throw new RuntimeException("Exception while running bundle 
finalizer", e);
-                  }
-                }));
-        getOutputBuilder().addFinalizeIds(id);
+      List<Pair<Instant, BundleFinalizer.Callback>> stepCallbacks =
+          stepContext.flushBundleFinalizerCallbacks();
+      if (!stepCallbacks.isEmpty()) {
+        Map<Long, Pair<Instant, Runnable>> targetMap = 
getOrCreateFinalizationCallbacks();
+        for (Pair<Instant, BundleFinalizer.Callback> bundleFinalizer : 
stepCallbacks) {
+          long id = ThreadLocalRandom.current().nextLong();
+          targetMap.put(
+              id,
+              Pair.of(
+                  bundleFinalizer.getLeft(),
+                  () -> {
+                    try {
+                      bundleFinalizer.getRight().onBundleSuccess();
+                    } catch (Exception e) {
+                      throw new RuntimeException("Exception while running 
bundle finalizer", e);
+                    }
+                  }));
+          getKeyOutputBuilder().addFinalizeIds(id);
+        }
       }
     }
 
     UnboundedReader<?> reader = activeReader;
     if (reader != null) {
-      Windmill.WorkItemCommitRequest.Builder builder = getOutputBuilder();
+      Windmill.WorkItemCommitRequest.Builder builder = getKeyOutputBuilder();
       Windmill.SourceState.Builder sourceStateBuilder = 
builder.getSourceStateUpdatesBuilder();
       final UnboundedSource.CheckpointMark checkpointMark = 
reader.getCheckpointMark();
       final Instant watermark = reader.getWatermark();
       long id = ThreadLocalRandom.current().nextLong();
       sourceStateBuilder.addFinalizeIds(id);
-      callbacks.put(
-          id,
-          Pair.of(
-              Instant.now().plus(Duration.standardMinutes(5)),
-              () -> {
-                try {
-                  checkpointMark.finalizeCheckpoint();
-                } catch (IOException e) {
-                  throw new RuntimeException("Exception while finalizing 
checkpoint", e);
-                }
-              }));
+      getOrCreateFinalizationCallbacks()
+          .put(
+              id,
+              Pair.of(
+                  Instant.now().plus(Duration.standardMinutes(5)),
+                  () -> {
+                    try {
+                      checkpointMark.finalizeCheckpoint();
+                    } catch (IOException e) {
+                      throw new RuntimeException("Exception while finalizing 
checkpoint", e);
+                    }
+                  }));
 
       @SuppressWarnings("unchecked")
       Coder<UnboundedSource.CheckpointMark> checkpointCoder =
@@ -692,19 +717,31 @@ public class StreamingModeExecutionContext
       // If activeReader is null, we might still have backlogBytes from an 
SDF. We ignore a reported
       // backlogBytes of 1 since older versions of the Java SDK use this value 
as a default when
       // RestrictionTracker.getProgress() or GetSize() are not defined.
-      getOutputBuilder().setSourceBacklogBytes(backlogBytes);
+      getKeyOutputBuilder().setSourceBacklogBytes(backlogBytes);
     }
 
-    this.finalizationCallbacks.putAll(callbacks);
-
-    getOutputBuilder()
+    getKeyOutputBuilder()
         
.setSourceBytesProcessed(computeSourceBytesProcessed(sourceBytesProcessCounterName));
 
     validateCommitRequestSize();
+
+    WorkItemCommitRequest workItemCommitRequest = 
getKeyOutputBuilder().build();
+    this.keyOutputBuilder = null;
+
+    if (multiKeyBundleOptions.multiKeyBundleEnabled()) {
+      if (this.workItemCommits == null) {
+        this.workItemCommits = new ArrayList<>();
+      }
+      this.workItemCommits.add(workItemCommitRequest);
+    } else {
+      checkState(this.workItemCommits == null);
+      this.workItemCommits = Collections.singletonList(workItemCommitRequest);
+    }
   }
 
   private void validateCommitRequestSize() {
-    Windmill.WorkItemCommitRequest.Builder currentBuilder = getOutputBuilder();
+    // TODO: Validate size of outputs at MultiKeyWorkItemCommitRequest level.
+    Windmill.WorkItemCommitRequest.Builder currentBuilder = 
getKeyOutputBuilder();
     Work currentWork = getWork();
     long byteLimit = operationalLimits.getMaxWorkItemCommitBytes();
     Windmill.WorkItemCommitRequest commitRequest = currentBuilder.build();
@@ -831,8 +868,7 @@ public class StreamingModeExecutionContext
     this.finishKeyCalled = false;
     this.computationKey = WindmillComputationKey.create(computationId, 
newWork.getShardedKey());
 
-    this.outputBuilder = createOutputBuilder(newWork);
-    this.outputBuilders.add(this.outputBuilder);
+    this.keyOutputBuilder = createOutputBuilder(newWork);
     newWork.setOnFailureListener(this.workBatchFailed);
 
     logHotKeyIfDetected(newWork, this.key);
@@ -860,23 +896,29 @@ public class StreamingModeExecutionContext
     }
   }
 
-  // Returns state bytes read during the bundle execution
-  public long getStateBytesRead() {
-    return stateBytesRead;
+  public void addBundleOutputMessages(Windmill.OutputMessageBundle 
outputBundle) {
+    if (this.bundleOutputMessages == null) {
+      this.bundleOutputMessages = new ArrayList<>();
+    }
+    this.bundleOutputMessages.add(outputBundle);
+  }
+
+  public void addBundlePubsubMessages(Windmill.PubSubMessageBundle 
pubsubBundle) {
+    if (this.bundlePubsubMessages == null) {
+      this.bundlePubsubMessages = new ArrayList<>();
+    }
+    this.bundlePubsubMessages.add(pubsubBundle);
   }
 
-  // Returns list of commit requests from the bundle
-  public List<Windmill.WorkItemCommitRequest> getWorkItemCommits() {
-    List<Windmill.WorkItemCommitRequest> commits = new 
ArrayList<>(outputBuilders.size());
-    for (Windmill.WorkItemCommitRequest.Builder builder : outputBuilders) {
-      commits.add(builder.build());
+  private Map<Long, Pair<Instant, Runnable>> 
getOrCreateFinalizationCallbacks() {
+    if (this.finalizationCallbacks == null) {
+      this.finalizationCallbacks = new HashMap<>();
     }
-    return commits;
+    return this.finalizationCallbacks;
   }
 
-  // Returns finalization callbacks recorded during the bundle execution
-  public Map<Long, Pair<Instant, Runnable>> getFinalizationCallbacks() {
-    return finalizationCallbacks;
+  public boolean multiKeyBundleEnabled() {
+    return multiKeyBundleOptions.multiKeyBundleEnabled();
   }
 
   // Returns the current key being processed or null if an unkeyed stage.
@@ -1241,7 +1283,7 @@ public class StreamingModeExecutionContext
 
     public void flushState() {
       if (stateFamily != null) {
-        WorkItemCommitRequest.Builder builder = getOutputBuilder();
+        WorkItemCommitRequest.Builder builder = getKeyOutputBuilder();
         checkStateNotNull(stateInternals).persist(builder);
         checkStateNotNull(systemTimerInternals).persistTo(builder);
         checkStateNotNull(userTimerInternals).persistTo(builder);
@@ -1456,7 +1498,7 @@ public class StreamingModeExecutionContext
               .setData(dataStream.toByteString())
               .setStateFamily(stateFamily);
 
-      getOutputBuilder().addGlobalDataUpdates(builder.build());
+      getKeyOutputBuilder().addGlobalDataUpdates(builder.build());
     }
 
     /** Fetch the given side input asynchronously and return true if it is 
present. */
@@ -1474,7 +1516,7 @@ public class StreamingModeExecutionContext
       String stateFamily = checkStateNotNull(this.stateFamily, "Tried to set 
global data request");
       sideInput =
           
Windmill.GlobalDataRequest.newBuilder(sideInput).setStateFamily(stateFamily).build();
-      WorkItemCommitRequest.Builder builder = getOutputBuilder();
+      WorkItemCommitRequest.Builder builder = getKeyOutputBuilder();
       builder.addGlobalDataRequests(sideInput);
       builder.addGlobalDataIdRequests(sideInput.getDataId());
     }
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WindmillSink.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WindmillSink.java
index 9d8a0f0da30..aef28a9bbec 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WindmillSink.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/WindmillSink.java
@@ -350,8 +350,7 @@ class WindmillSink<T> extends Sink<WindowedValue<T>> {
       return (long) key.size() + value.size() + metadata.size() + id.size() + 
offsetSize;
     }
 
-    @Override
-    public void close() throws IOException {
+    private void flush(boolean bundleLevel) {
       try {
         outputBuilder.setDestinationStreamId(destinationName);
 
@@ -359,17 +358,35 @@ class WindmillSink<T> extends Sink<WindowedValue<T>> {
           outputBuilder.addBundles(keyedOutput.build());
         }
         if (outputBuilder.getBundlesCount() > 0) {
-          context.getOutputBuilder().addOutputMessages(outputBuilder.build());
+          Windmill.OutputMessageBundle bundle = outputBuilder.build();
+          if (bundleLevel) {
+            context.addBundleOutputMessages(bundle);
+          } else {
+            context.getKeyOutputBuilder().addOutputMessages(bundle);
+          }
         }
       } finally {
         outputBuilder.clear();
+        productionMap.clear();
       }
-      productionMap.clear();
+    }
+
+    @Override
+    public void finishKey(@Nullable Object key) throws IOException {
+      if (context.multiKeyBundleEnabled()) {
+        flush(/* bundleLevel= */ false);
+      }
+    }
+
+    @Override
+    public void close() throws IOException {
+      flush(/* bundleLevel= */ context.multiKeyBundleEnabled());
     }
 
     @Override
     public void abort() throws IOException {
-      close();
+      outputBuilder.clear();
+      productionMap.clear();
     }
   }
 
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/Sink.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/Sink.java
index 7bf9ec99b6a..36fde605a13 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/Sink.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/Sink.java
@@ -18,6 +18,7 @@
 package org.apache.beam.runners.dataflow.worker.util.common.worker;
 
 import java.io.IOException;
+import org.checkerframework.checker.nullness.qual.Nullable;
 
 /**
  * Abstract base class for Sinks.
@@ -36,6 +37,11 @@ public abstract class Sink<T> {
     /** Adds a value to the sink. Returns the size in bytes of the data 
written. */
     public long add(ElemT value) throws IOException;
 
+    /**
+     * Called when all elements for a specific key have been processed. Called 
only for Streaming
+     */
+    public default void finishKey(@Nullable Object key) throws IOException {}
+
     /**
      * {@inheritDoc}
      *
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperation.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperation.java
index a97c9920b9a..013913d97c5 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperation.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperation.java
@@ -107,7 +107,14 @@ public class WriteOperation extends ReceivingOperation {
   }
 
   @Override
-  public void finishKey(@Nullable Object key) throws Exception {}
+  public void finishKey(@Nullable Object key) throws Exception {
+    try (Closeable scope = context.enterProcess()) {
+      checkStarted();
+      if (writer != null) {
+        writer.finishKey(key);
+      }
+    }
+  }
 
   @Override
   public void abort() throws Exception {
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ExecuteWorkResult.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ExecuteWorkResult.java
new file mode 100644
index 00000000000..dad9e7134d0
--- /dev/null
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/ExecuteWorkResult.java
@@ -0,0 +1,57 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.beam.runners.dataflow.worker.windmill.work.processing;
+
+import com.google.auto.value.AutoValue;
+import java.util.List;
+import java.util.Map;
+import javax.annotation.concurrent.Immutable;
+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.sdk.annotations.Internal;
+import org.joda.time.Instant;
+
+@Internal
+@Immutable
+@AutoValue
+public abstract class ExecuteWorkResult {
+  public static ExecuteWorkResult create(
+      List<Windmill.WorkItemCommitRequest> workItemCommits,
+      List<Windmill.OutputMessageBundle> bundleOutputMessages,
+      List<Windmill.PubSubMessageBundle> bundlePubsubMessages,
+      Map<Long, Pair<Instant, Runnable>> finalizationCallbacks,
+      long stateBytesRead) {
+    return new AutoValue_ExecuteWorkResult(
+        workItemCommits,
+        bundleOutputMessages,
+        bundlePubsubMessages,
+        finalizationCallbacks,
+        stateBytesRead);
+  }
+
+  public abstract List<Windmill.WorkItemCommitRequest> workItemCommits();
+
+  public abstract List<Windmill.OutputMessageBundle> bundleOutputMessages();
+
+  public abstract List<Windmill.PubSubMessageBundle> bundlePubsubMessages();
+
+  // Map<finalizerId, Pair<callbackExpiration, callback>>
+  public abstract Map<Long, Pair<Instant, Runnable>> finalizationCallbacks();
+
+  public abstract long stateBytesRead();
+}
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java
 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java
index 78e429d4089..aff42b8468f 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java
@@ -20,17 +20,14 @@ package 
org.apache.beam.runners.dataflow.worker.windmill.work.processing;
 import static 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkState;
 
 import com.google.api.services.dataflow.model.MapTask;
-import com.google.auto.value.AutoValue;
 import java.util.ArrayList;
 import java.util.List;
-import java.util.Map;
 import java.util.concurrent.ConcurrentMap;
 import java.util.concurrent.ScheduledExecutorService;
 import java.util.concurrent.TimeUnit;
 import java.util.function.Function;
 import java.util.function.Supplier;
 import javax.annotation.concurrent.ThreadSafe;
-import org.apache.beam.repackaged.core.org.apache.commons.lang3.tuple.Pair;
 import org.apache.beam.runners.dataflow.options.DataflowWorkerHarnessOptions;
 import org.apache.beam.runners.dataflow.worker.DataflowExecutionStateSampler;
 import org.apache.beam.runners.dataflow.worker.DataflowMapTaskExecutorFactory;
@@ -252,10 +249,15 @@ public class StreamingWorkScheduler {
           executeWork(work, stageInfo, computationState, handle, 
keyTransitionListener);
       List<Work> workBatch = handle.getWorkBatch();
       List<Windmill.WorkItemCommitRequest> workItemCommits = 
executeWorkResult.workItemCommits();
+      List<Windmill.OutputMessageBundle> bundleOutputMessages =
+          executeWorkResult.bundleOutputMessages();
+      List<Windmill.PubSubMessageBundle> bundlePubsubMessages =
+          executeWorkResult.bundlePubsubMessages();
 
       
commitFinalizer.cacheCommitFinalizers(executeWorkResult.finalizationCallbacks());
 
-      commitWorkBatch(computationState, workBatch, workItemCommits);
+      commitWorkBatch(
+          computationState, workBatch, workItemCommits, bundleOutputMessages, 
bundlePubsubMessages);
 
       recordProcessingStats(workBatch, workItemCommits, 
executeWorkResult.stateBytesRead());
       LOG.debug("Processing done for work batch size: {}", workBatch.size());
@@ -326,26 +328,17 @@ public class StreamingWorkScheduler {
       computationWorkExecutor.executeWork(
           work, workExecutor, handle, keyTransitionListener, 
onFailedWorkHandler);
 
-      List<Windmill.WorkItemCommitRequest> workItemCommits;
-      Map<Long, Pair<Instant, Runnable>> finalizationCallbacks;
-      long stateBytesRead;
-      {
-        if (context.workIsFailed()) {
-          throw new 
WorkItemCancelledException(work.getWorkItem().getShardingKey());
-        }
-        context.flushState();
-
-        workItemCommits = context.getWorkItemCommits();
-        finalizationCallbacks = context.getFinalizationCallbacks();
-        stateBytesRead = context.getStateBytesRead();
-
-        context.reset(); // Don't use context after this.
+      if (context.workIsFailed()) {
+        throw new 
WorkItemCancelledException(work.getWorkItem().getShardingKey());
       }
+      // Don't use context after this.
+      ExecuteWorkResult executeWorkResult = context.flushStateAndReset();
+
       // Release the execution state for another thread to use.
       computationState.releaseComputationWorkExecutor(computationWorkExecutor);
       computationWorkExecutor = null;
 
-      return ExecuteWorkResult.create(workItemCommits, finalizationCallbacks, 
stateBytesRead);
+      return executeWorkResult;
     } catch (Throwable t) {
       if (computationWorkExecutor != null) {
         // If processing failed due to a thrown exception, close the 
executionState. Do not
@@ -380,13 +373,22 @@ public class StreamingWorkScheduler {
   private void commitWorkBatch(
       ComputationState computationState,
       List<Work> workBatch,
-      List<Windmill.WorkItemCommitRequest> workItemCommits) {
+      List<Windmill.WorkItemCommitRequest> workItemCommits,
+      List<Windmill.OutputMessageBundle> bundleOutputMessages,
+      List<Windmill.PubSubMessageBundle> bundlePubsubMessages) {
     if (workBatch.isEmpty()) {
       return;
     }
     if (workBatch.size() > 1 || multiKeyBundleOptions.multiKeyBundleEnabled()) 
{
-      commitMultiKeyWorkBatch(computationState, workBatch, workItemCommits);
+      commitMultiKeyWorkBatch(
+          computationState, workBatch, workItemCommits, bundleOutputMessages, 
bundlePubsubMessages);
     } else {
+      checkState(
+          bundleOutputMessages.isEmpty(),
+          "bundleOutputMessages should be empty when calling 
commitSingleKeyWork");
+      checkState(
+          bundlePubsubMessages.isEmpty(),
+          "bundlePubsubMessages should be empty when calling 
commitSingleKeyWork");
       commitSingleKeyWork(computationState, workBatch.get(0), 
workItemCommits.get(0));
     }
   }
@@ -394,11 +396,19 @@ public class StreamingWorkScheduler {
   private void commitMultiKeyWorkBatch(
       ComputationState computationState,
       List<Work> workBatch,
-      List<Windmill.WorkItemCommitRequest> workItemCommits) {
+      List<Windmill.WorkItemCommitRequest> workItemCommits,
+      List<Windmill.OutputMessageBundle> bundleOutputMessages,
+      List<Windmill.PubSubMessageBundle> bundlePubsubMessages) {
     checkState(!workBatch.isEmpty());
     checkState(workBatch.size() == workItemCommits.size());
     Windmill.MultiKeyWorkItemCommitRequest.Builder multiKeyBuilder =
         Windmill.MultiKeyWorkItemCommitRequest.newBuilder();
+    if (!bundleOutputMessages.isEmpty()) {
+      multiKeyBuilder.addAllOutputMessages(bundleOutputMessages);
+    }
+    if (!bundlePubsubMessages.isEmpty()) {
+      multiKeyBuilder.addAllPubsubMessages(bundlePubsubMessages);
+    }
 
     Work primaryWork = workBatch.get(0);
     Work.KeyGroup keyGroup = primaryWork.getKeyGroup();
@@ -506,22 +516,4 @@ public class StreamingWorkScheduler {
       }
     };
   }
-
-  @AutoValue
-  abstract static class ExecuteWorkResult {
-    static ExecuteWorkResult create(
-        List<Windmill.WorkItemCommitRequest> workItemCommits,
-        Map<Long, Pair<Instant, Runnable>> finalizationCallbacks,
-        long stateBytesRead) {
-      return new AutoValue_StreamingWorkScheduler_ExecuteWorkResult(
-          workItemCommits, finalizationCallbacks, stateBytesRead);
-    }
-
-    abstract List<Windmill.WorkItemCommitRequest> workItemCommits();
-
-    // Map<finalizerId, Pair<callbackExpiration, callback>>
-    abstract Map<Long, Pair<Instant, Runnable>> finalizationCallbacks();
-
-    abstract long stateBytesRead();
-  }
 }
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSinkTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSinkTest.java
index d8822ce4937..dfb4f221726 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSinkTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubDynamicSinkTest.java
@@ -18,6 +18,7 @@
 package org.apache.beam.runners.dataflow.worker;
 
 import static org.junit.Assert.assertEquals;
+import static org.mockito.Mockito.verify;
 import static org.mockito.Mockito.when;
 
 import java.nio.charset.StandardCharsets;
@@ -41,6 +42,7 @@ import org.junit.Before;
 import org.junit.Test;
 import org.junit.runner.RunWith;
 import org.junit.runners.JUnit4;
+import org.mockito.ArgumentCaptor;
 import org.mockito.Mock;
 import org.mockito.MockitoAnnotations;
 
@@ -61,7 +63,7 @@ public class PubsubDynamicSinkTest {
             .setKey(ByteString.copyFromUtf8("key"))
             .setWorkToken(0);
 
-    when(mockContext.getOutputBuilder()).thenReturn(outputBuilder);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
 
     Map<String, Object> spec = new HashMap<>();
     spec.put(PropertyNames.OBJECT_TYPE_NAME, "PubsubDynamicSink");
@@ -161,4 +163,210 @@ public class PubsubDynamicSinkTest {
             .build();
     assertEquals(expectedCommit, outputBuilder.build());
   }
+
+  @Test
+  public void testSingleKey_finishKeyDoesNotFlush_closeAttachesToKey() throws 
Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(false);
+
+    Windmill.WorkItemCommitRequest.Builder outputBuilder =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key"))
+            .setWorkToken(0);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "PubsubDynamicSink");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubDynamicSink sink =
+        (PubsubDynamicSink)
+            SinkRegistry.defaultRegistry()
+                .create(
+                    cloudSinkSpec,
+                    WindowedValues.getFullCoder(VoidCoder.of(), 
IntervalWindow.getCoder()),
+                    null,
+                    mockContext,
+                    null)
+                .getUnderlyingSink();
+
+    Sink.SinkWriter<WindowedValue<PubsubMessage>> writer = sink.writer();
+    byte[] payload0 = "msg0".getBytes(StandardCharsets.UTF_8);
+    byte[] payload1 = "msg1".getBytes(StandardCharsets.UTF_8);
+
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage(payload0, null).withTopic("topic1"), new 
Instant(0)));
+
+    // In single-key mode, finishKey does not flush
+    writer.finishKey("key");
+    assertEquals(0, outputBuilder.getPubsubMessagesCount());
+
+    // close flushes all outputs into the key's outputBuilder
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage(payload1, null).withTopic("topic2"), new 
Instant(1000)));
+    writer.close();
+
+    assertEquals(2, outputBuilder.getPubsubMessagesCount());
+    Map<String, Windmill.PubSubMessageBundle> bundlesByTopic = new HashMap<>();
+    for (Windmill.PubSubMessageBundle bundle : 
outputBuilder.getPubsubMessagesList()) {
+      bundlesByTopic.put(bundle.getTopic(), bundle);
+    }
+    assertEquals(1, bundlesByTopic.get("topic1").getMessagesCount());
+    assertEquals(1, bundlesByTopic.get("topic2").getMessagesCount());
+    Pubsub.PubsubMessage pubsubMsg0 =
+        
Pubsub.PubsubMessage.parseFrom(bundlesByTopic.get("topic1").getMessages(0).getData());
+    assertEquals(ByteString.copyFrom(payload0), pubsubMsg0.getData());
+    Pubsub.PubsubMessage pubsubMsg1 =
+        
Pubsub.PubsubMessage.parseFrom(bundlesByTopic.get("topic2").getMessages(0).getData());
+    assertEquals(ByteString.copyFrom(payload1), pubsubMsg1.getData());
+  }
+
+  @Test
+  public void testMultiKey_flushesAllTopicsToBundleLevelAtClose() throws 
Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(true);
+
+    Windmill.WorkItemCommitRequest.Builder outputBuilderKey1 =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1);
+    Windmill.WorkItemCommitRequest.Builder outputBuilderKey2 =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key2"))
+            .setWorkToken(2);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilderKey1);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "PubsubDynamicSink");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubDynamicSink sink =
+        (PubsubDynamicSink)
+            SinkRegistry.defaultRegistry()
+                .create(
+                    cloudSinkSpec,
+                    WindowedValues.getFullCoder(VoidCoder.of(), 
IntervalWindow.getCoder()),
+                    null,
+                    mockContext,
+                    null)
+                .getUnderlyingSink();
+
+    Sink.SinkWriter<WindowedValue<PubsubMessage>> writer = sink.writer();
+
+    // 1. Process Key 1 messages (to topicA and topicB)
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("k1-msgA".getBytes(StandardCharsets.UTF_8), 
null).withTopic("topicA"),
+            new Instant(0)));
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("k1-msgB".getBytes(StandardCharsets.UTF_8), 
null).withTopic("topicB"),
+            new Instant(10)));
+    writer.finishKey("key1");
+
+    // In multi-key mode, finishKey does not flush to key-level commit
+    assertEquals(0, outputBuilderKey1.getPubsubMessagesCount());
+
+    // 2. Process Key 2 messages (to topicB and topicC)
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilderKey2);
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("k2-msgB".getBytes(StandardCharsets.UTF_8), 
null).withTopic("topicB"),
+            new Instant(100)));
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("k2-msgC".getBytes(StandardCharsets.UTF_8), 
null).withTopic("topicC"),
+            new Instant(110)));
+    writer.finishKey("key2");
+
+    // In multi-key mode, finishKey does not flush to key-level commit
+    assertEquals(0, outputBuilderKey2.getPubsubMessagesCount());
+
+    // 3. Process finishBundle messages (to topicC) and close
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("bundle-tC".getBytes(StandardCharsets.UTF_8), 
null)
+                .withTopic("topicC"),
+            new Instant(200)));
+    writer.close();
+
+    // Verify Bundle-level flush expectations: all messages grouped by topic 
at bundle level
+    ArgumentCaptor<Windmill.PubSubMessageBundle> captor =
+        ArgumentCaptor.forClass(Windmill.PubSubMessageBundle.class);
+    verify(mockContext, 
org.mockito.Mockito.times(3)).addBundlePubsubMessages(captor.capture());
+    Map<String, Integer> actualBundleTopicCounts = new HashMap<>();
+    for (Windmill.PubSubMessageBundle b : captor.getAllValues()) {
+      actualBundleTopicCounts.put(b.getTopic(), b.getMessagesCount());
+    }
+    assertEquals(Map.of("topicA", 1, "topicB", 2, "topicC", 2), 
actualBundleTopicCounts);
+  }
+
+  @Test
+  public void testMultiKey_emptyBundleFlushesNothing() throws Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(true);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "PubsubDynamicSink");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubDynamicSink sink =
+        (PubsubDynamicSink)
+            SinkRegistry.defaultRegistry()
+                .create(
+                    cloudSinkSpec,
+                    WindowedValues.getFullCoder(VoidCoder.of(), 
IntervalWindow.getCoder()),
+                    null,
+                    mockContext,
+                    null)
+                .getUnderlyingSink();
+
+    Sink.SinkWriter<WindowedValue<PubsubMessage>> writer = sink.writer();
+    writer.finishKey("key1");
+    writer.close();
+
+    verify(mockContext, org.mockito.Mockito.never())
+        .addBundlePubsubMessages(org.mockito.ArgumentMatchers.any());
+  }
+
+  @Test
+  public void testAbort() throws Exception {
+    Windmill.WorkItemCommitRequest.Builder outputBuilder =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key"))
+            .setWorkToken(0);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "PubsubDynamicSink");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubDynamicSink sink =
+        (PubsubDynamicSink)
+            SinkRegistry.defaultRegistry()
+                .create(
+                    cloudSinkSpec,
+                    WindowedValues.getFullCoder(VoidCoder.of(), 
IntervalWindow.getCoder()),
+                    null,
+                    mockContext,
+                    null)
+                .getUnderlyingSink();
+
+    Sink.SinkWriter<WindowedValue<PubsubMessage>> writer = sink.writer();
+
+    // Buffer and abort
+    writer.add(
+        WindowedValues.timestampedValueInGlobalWindow(
+            new PubsubMessage("aborted".getBytes(StandardCharsets.UTF_8), 
null).withTopic("topic1"),
+            new Instant(0)));
+    writer.abort();
+    assertEquals(0, outputBuilder.getPubsubMessagesCount());
+  }
 }
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubSinkTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubSinkTest.java
index 5327cd17241..6b9c911715b 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubSinkTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/PubsubSinkTest.java
@@ -19,6 +19,7 @@ package org.apache.beam.runners.dataflow.worker;
 
 import static org.junit.Assert.assertEquals;
 import static org.junit.Assert.assertThrows;
+import static org.mockito.Mockito.verify;
 import static org.mockito.Mockito.when;
 
 import java.io.IOException;
@@ -43,6 +44,7 @@ import org.junit.Before;
 import org.junit.Test;
 import org.junit.runner.RunWith;
 import org.junit.runners.JUnit4;
+import org.mockito.ArgumentCaptor;
 import org.mockito.Mock;
 import org.mockito.MockitoAnnotations;
 
@@ -62,7 +64,7 @@ public class PubsubSinkTest {
             .setKey(ByteString.copyFromUtf8("key"))
             .setWorkToken(0);
 
-    when(mockContext.getOutputBuilder()).thenReturn(outputBuilder);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
 
     Map<String, Object> spec = new HashMap<>();
     spec.put(PropertyNames.OBJECT_TYPE_NAME, "");
@@ -180,4 +182,174 @@ public class PubsubSinkTest {
         CoderException.class,
         () -> writer.add(WindowedValues.timestampedValueInGlobalWindow("e0", 
new Instant(0))));
   }
+
+  @Test
+  public void testSingleKey_finishKeyDoesNotFlush_closeAttachesToKey() throws 
Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(false);
+
+    Windmill.WorkItemCommitRequest.Builder outputBuilder =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key"))
+            .setWorkToken(0);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "");
+    spec.put(PropertyNames.PUBSUB_TOPIC, "topic");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubSink.Factory factory = new PubsubSink.Factory();
+    PubsubSink<String> sink =
+        (PubsubSink<String>)
+            factory.create(
+                cloudSinkSpec,
+                WindowedValues.getFullCoder(StringUtf8Coder.of(), 
IntervalWindow.getCoder()),
+                null,
+                mockContext,
+                null);
+
+    Sink.SinkWriter<WindowedValue<String>> writer = sink.writer();
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("e0", new 
Instant(0)));
+
+    // In single key mode, finishKey should not flush
+    writer.finishKey("key");
+    assertEquals(0, outputBuilder.getPubsubMessagesCount());
+
+    // close should flush and attach to the key's outputBuilder
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("e1", new 
Instant(1000)));
+    writer.close();
+
+    assertEquals(1, outputBuilder.getPubsubMessagesCount());
+    Windmill.PubSubMessageBundle bundle = outputBuilder.getPubsubMessages(0);
+    assertEquals("topic", bundle.getTopic());
+    assertEquals(2, bundle.getMessagesCount());
+    assertEquals("e0", bundle.getMessages(0).getData().toStringUtf8());
+    assertEquals("e1", bundle.getMessages(1).getData().toStringUtf8());
+  }
+
+  @Test
+  public void testMultiKey_flushesAllMessagesToBundleLevelAtClose() throws 
Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(true);
+
+    Windmill.WorkItemCommitRequest.Builder outputBuilderKey1 =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1);
+    Windmill.WorkItemCommitRequest.Builder outputBuilderKey2 =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key2"))
+            .setWorkToken(2);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilderKey1);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "");
+    spec.put(PropertyNames.PUBSUB_TOPIC, "topic");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubSink.Factory factory = new PubsubSink.Factory();
+    PubsubSink<String> sink =
+        (PubsubSink<String>)
+            factory.create(
+                cloudSinkSpec,
+                WindowedValues.getFullCoder(StringUtf8Coder.of(), 
IntervalWindow.getCoder()),
+                null,
+                mockContext,
+                null);
+
+    Sink.SinkWriter<WindowedValue<String>> writer = sink.writer();
+
+    // 1. Process Key 1 messages
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("k1-msg1", new 
Instant(0)));
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("k1-msg2", new 
Instant(10)));
+    writer.finishKey("key1");
+
+    // In multi-key mode, finishKey does not flush to key-level commit
+    assertEquals(0, outputBuilderKey1.getPubsubMessagesCount());
+
+    // 2. Process Key 2 messages
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilderKey2);
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("k2-msg1", new 
Instant(100)));
+    writer.finishKey("key2");
+
+    // In multi-key mode, finishKey does not flush to key-level commit
+    assertEquals(0, outputBuilderKey2.getPubsubMessagesCount());
+
+    // 3. Process finishBundle messages and close
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("bundle-msg", new 
Instant(200)));
+    writer.close();
+
+    // Verify all messages across keys and finishBundle flush to bundle level 
at close
+    ArgumentCaptor<Windmill.PubSubMessageBundle> captor =
+        ArgumentCaptor.forClass(Windmill.PubSubMessageBundle.class);
+    verify(mockContext).addBundlePubsubMessages(captor.capture());
+    Windmill.PubSubMessageBundle bundleLevel = captor.getValue();
+    assertEquals("topic", bundleLevel.getTopic());
+    assertEquals(4, bundleLevel.getMessagesCount());
+    assertEquals("k1-msg1", 
bundleLevel.getMessages(0).getData().toStringUtf8());
+    assertEquals("k1-msg2", 
bundleLevel.getMessages(1).getData().toStringUtf8());
+    assertEquals("k2-msg1", 
bundleLevel.getMessages(2).getData().toStringUtf8());
+    assertEquals("bundle-msg", 
bundleLevel.getMessages(3).getData().toStringUtf8());
+  }
+
+  @Test
+  public void testMultiKey_emptyBundleFlushesNothing() throws Exception {
+    when(mockContext.multiKeyBundleEnabled()).thenReturn(true);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "");
+    spec.put(PropertyNames.PUBSUB_TOPIC, "topic");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubSink.Factory factory = new PubsubSink.Factory();
+    PubsubSink<String> sink =
+        (PubsubSink<String>)
+            factory.create(
+                cloudSinkSpec,
+                WindowedValues.getFullCoder(StringUtf8Coder.of(), 
IntervalWindow.getCoder()),
+                null,
+                mockContext,
+                null);
+
+    Sink.SinkWriter<WindowedValue<String>> writer = sink.writer();
+    writer.finishKey("key1");
+    writer.close();
+
+    verify(mockContext, org.mockito.Mockito.never())
+        .addBundlePubsubMessages(org.mockito.ArgumentMatchers.any());
+  }
+
+  @Test
+  public void testAbort() throws Exception {
+    Windmill.WorkItemCommitRequest.Builder outputBuilder =
+        Windmill.WorkItemCommitRequest.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key"))
+            .setWorkToken(0);
+    when(mockContext.getKeyOutputBuilder()).thenReturn(outputBuilder);
+
+    Map<String, Object> spec = new HashMap<>();
+    spec.put(PropertyNames.OBJECT_TYPE_NAME, "");
+    spec.put(PropertyNames.PUBSUB_TOPIC, "topic");
+    spec.put(PropertyNames.PUBSUB_TIMESTAMP_ATTRIBUTE, "ts");
+    spec.put(PropertyNames.PUBSUB_ID_ATTRIBUTE, "id");
+    CloudObject cloudSinkSpec = CloudObject.fromSpec(spec);
+    PubsubSink.Factory factory = new PubsubSink.Factory();
+    PubsubSink<String> sink =
+        (PubsubSink<String>)
+            factory.create(
+                cloudSinkSpec,
+                WindowedValues.getFullCoder(StringUtf8Coder.of(), 
IntervalWindow.getCoder()),
+                null,
+                mockContext,
+                null);
+
+    Sink.SinkWriter<WindowedValue<String>> writer = sink.writer();
+
+    // Buffer message and abort
+    writer.add(WindowedValues.timestampedValueInGlobalWindow("msg-aborted", 
new Instant(0)));
+    writer.abort();
+    assertEquals(0, outputBuilder.getPubsubMessagesCount());
+  }
 }
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java
index c0206eed17a..6fc19c30324 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java
@@ -33,6 +33,7 @@ import static org.junit.Assert.assertNull;
 import static org.junit.Assert.assertSame;
 import static org.junit.Assert.assertTrue;
 import static org.junit.Assert.fail;
+import static org.junit.Assume.assumeTrue;
 import static org.mockito.ArgumentMatchers.any;
 import static org.mockito.ArgumentMatchers.eq;
 import static org.mockito.ArgumentMatchers.nullable;
@@ -1659,6 +1660,443 @@ public class StreamingDataflowWorkerTest {
     worker.stop();
   }
 
+  private void runMultiKeyCombinationTest(
+      Map<String, List<String>> processOutputs, List<KV<String, String>> 
finishBundleOutputs)
+      throws Exception {
+    assumeTrue("Multi-key bundling is only supported in Streaming Engine", 
streamingEngine);
+    server.clearCommitsReceived();
+    StreamingDataflowWorker worker =
+        makeMultiKeyEnabledWorker(
+            new ConfigurableMultiKeyDoFn(processOutputs, finishBundleOutputs));
+    worker.start();
+
+    String batchInputText =
+        "work {"
+            + "  computation_id: \""
+            + DEFAULT_COMPUTATION_ID
+            + "\""
+            + "  input_data_watermark: 0"
+            + "  work {"
+            + "    key: \"key1\""
+            + "    sharding_key: 1"
+            + "    work_token: 1"
+            + "    cache_token: 2"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data1\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "  work {"
+            + "    key: \"key2\""
+            + "    sharding_key: 2"
+            + "    work_token: 2"
+            + "    cache_token: 3"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data2\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "}";
+    Windmill.GetWorkResponse batchInput =
+        buildInput(
+            batchInputText,
+            CoderUtils.encodeToByteArray(
+                CollectionCoder.of(IntervalWindow.getCoder()),
+                Collections.singletonList(DEFAULT_WINDOW)));
+
+    
server.whenGetDataCalled().answerByDefault(StreamingDataflowWorkerTest::emptyDataResponder);
+    server.whenGetWorkCalled().thenReturn(batchInput);
+
+    Map<Long, Windmill.WorkItemCommitRequest> result = 
server.waitForAndGetCommits(2);
+    assertEquals(2, result.size());
+
+    // Verify Key 1 commit: should only contain process outputs for key1
+    assertTrue(result.containsKey(1L));
+    Windmill.WorkItemCommitRequest commit1 = result.get(1L);
+    assertEquals("key1", commit1.getKey().toStringUtf8());
+    List<String> expectedKey1Outputs = processOutputs.getOrDefault("key1", 
Collections.emptyList());
+    if (expectedKey1Outputs.isEmpty()) {
+      assertEquals(0, commit1.getOutputMessagesCount());
+    } else {
+      assertEquals(1, commit1.getOutputMessagesCount());
+      Windmill.OutputMessageBundle outputBundle1 = 
commit1.getOutputMessages(0);
+      assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
outputBundle1.getDestinationStreamId());
+      assertEquals(1, outputBundle1.getBundlesCount());
+      Windmill.KeyedMessageBundle keyedBundle1 = outputBundle1.getBundles(0);
+      assertEquals("key1", keyedBundle1.getKey().toStringUtf8());
+      assertEquals(expectedKey1Outputs.size(), 
keyedBundle1.getMessagesCount());
+      for (int i = 0; i < expectedKey1Outputs.size(); i++) {
+        assertEquals(
+            expectedKey1Outputs.get(i), 
keyedBundle1.getMessages(i).getData().toStringUtf8());
+      }
+    }
+    // Verify Key 2 commit: should only contain process outputs for key2 (NOT 
finishBundle outputs)
+    assertTrue(result.containsKey(2L));
+    Windmill.WorkItemCommitRequest commit2 = result.get(2L);
+    assertEquals("key2", commit2.getKey().toStringUtf8());
+    List<String> expectedKey2Outputs = processOutputs.getOrDefault("key2", 
Collections.emptyList());
+    if (expectedKey2Outputs.isEmpty()) {
+      assertEquals(0, commit2.getOutputMessagesCount());
+    } else {
+      assertEquals(1, commit2.getOutputMessagesCount());
+      Windmill.OutputMessageBundle outputBundle2 = 
commit2.getOutputMessages(0);
+      assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
outputBundle2.getDestinationStreamId());
+      assertEquals(1, outputBundle2.getBundlesCount());
+      Windmill.KeyedMessageBundle keyedBundle2 = outputBundle2.getBundles(0);
+      assertEquals("key2", keyedBundle2.getKey().toStringUtf8());
+      assertEquals(expectedKey2Outputs.size(), 
keyedBundle2.getMessagesCount());
+      for (int i = 0; i < expectedKey2Outputs.size(); i++) {
+        assertEquals(
+            expectedKey2Outputs.get(i), 
keyedBundle2.getMessages(i).getData().toStringUtf8());
+      }
+    }
+
+    // Verify MultiKey commit: should contain all finishBundle outputs at the 
bundle level
+    List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+        server.getMultiKeyCommitsReceived();
+    assertEquals(1, multiKeyCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = 
multiKeyCommits.get(0);
+    if (finishBundleOutputs.isEmpty()) {
+      assertEquals(0, multiKeyCommit.getOutputMessagesCount());
+    } else {
+      assertEquals(1, multiKeyCommit.getOutputMessagesCount());
+      Windmill.OutputMessageBundle outputBundle_fb = 
multiKeyCommit.getOutputMessages(0);
+      assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
outputBundle_fb.getDestinationStreamId());
+      Map<String, List<String>> expectedFbByKey = new HashMap<>();
+      for (KV<String, String> kv : finishBundleOutputs) {
+        expectedFbByKey.computeIfAbsent(kv.getKey(), k -> new 
ArrayList<>()).add(kv.getValue());
+      }
+      assertEquals(expectedFbByKey.size(), outputBundle_fb.getBundlesCount());
+      for (Windmill.KeyedMessageBundle keyedBundle : 
outputBundle_fb.getBundlesList()) {
+        String key = keyedBundle.getKey().toStringUtf8();
+        assertTrue(expectedFbByKey.containsKey(key));
+        List<String> expectedValues = expectedFbByKey.get(key);
+        assertEquals(expectedValues.size(), keyedBundle.getMessagesCount());
+        for (int i = 0; i < expectedValues.size(); i++) {
+          assertEquals(expectedValues.get(i), 
keyedBundle.getMessages(i).getData().toStringUtf8());
+        }
+      }
+    }
+
+    worker.stop();
+  }
+
+  @Test
+  public void testMultiKey_allCombinationsOfProcessAndFinishBundleOutputs() 
throws Exception {
+    if (!streamingEngine) {
+      return;
+    }
+    List<List<String>> key1Options =
+        List.of(Collections.emptyList(), List.of("k1_out1"), 
List.of("k1_out1", "k1_out2"));
+
+    List<List<String>> key2Options =
+        List.of(Collections.emptyList(), List.of("k2_out1"), 
List.of("k2_out1", "k2_out2"));
+
+    List<List<KV<String, String>>> finishBundleOptions =
+        List.of(
+            Collections.emptyList(),
+            List.of(KV.of("fb_key", "fb_val1")),
+            List.of(KV.of("fb_key1", "fb_val1"), KV.of("fb_key2", "fb_val2")));
+
+    for (List<String> k1Out : key1Options) {
+      for (List<String> k2Out : key2Options) {
+        for (List<KV<String, String>> fbOut : finishBundleOptions) {
+          Map<String, List<String>> processOutputs = new HashMap<>();
+          if (!k1Out.isEmpty()) {
+            processOutputs.put("key1", k1Out);
+          }
+          if (!k2Out.isEmpty()) {
+            processOutputs.put("key2", k2Out);
+          }
+          runMultiKeyCombinationTest(processOutputs, fbOut);
+        }
+      }
+    }
+  }
+
+  @Test
+  public void testSingleKey_processAndFinishBundleOutputsAttachedToSameKey() 
throws Exception {
+    KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
+    List<ParallelInstruction> instructions =
+        Arrays.asList(
+            makeSourceInstruction(kvCoder),
+            makeDoFnInstruction(
+                new ConfigurableMultiKeyDoFn(
+                    ImmutableMap.of("key1", ImmutableList.of("data1")),
+                    ImmutableList.of(KV.of("finish_key", "finish_value"))),
+                0,
+                kvCoder),
+            makeSinkInstruction(kvCoder, 1));
+
+    StreamingDataflowWorker worker =
+        
makeWorker(defaultWorkerParams().setInstructions(instructions).build());
+    worker.start();
+
+    String input =
+        "work {"
+            + "  computation_id: \""
+            + DEFAULT_COMPUTATION_ID
+            + "\""
+            + "  input_data_watermark: 0"
+            + "  work {"
+            + "    key: \"key1\""
+            + "    sharding_key: 1"
+            + "    work_token: 1"
+            + "    cache_token: 2"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data1\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "}";
+    Windmill.GetWorkResponse workResponse =
+        buildInput(
+            input,
+            CoderUtils.encodeToByteArray(
+                CollectionCoder.of(IntervalWindow.getCoder()),
+                Collections.singletonList(DEFAULT_WINDOW)));
+
+    server.clearCommitsReceived();
+    
server.whenGetDataCalled().answerByDefault(StreamingDataflowWorkerTest::emptyDataResponder);
+    server.whenGetWorkCalled().thenReturn(workResponse);
+
+    Map<Long, Windmill.WorkItemCommitRequest> result = 
server.waitForAndGetCommits(1);
+    assertEquals(1, result.size());
+
+    assertTrue(result.containsKey(1L));
+    Windmill.WorkItemCommitRequest commit = result.get(1L);
+    assertEquals("key1", commit.getKey().toStringUtf8());
+    // In single-key mode, finishKey does not flush; close flushes all outputs 
into the single key's
+    // commit
+    assertEquals(1, commit.getOutputMessagesCount());
+    Windmill.OutputMessageBundle outputBundle = commit.getOutputMessages(0);
+    assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
outputBundle.getDestinationStreamId());
+    assertEquals(2, outputBundle.getBundlesCount());
+    Map<String, String> outputsByKey = new HashMap<>();
+    for (Windmill.KeyedMessageBundle bundle : outputBundle.getBundlesList()) {
+      assertEquals(1, bundle.getMessagesCount());
+      outputsByKey.put(
+          bundle.getKey().toStringUtf8(), 
bundle.getMessages(0).getData().toStringUtf8());
+    }
+    assertEquals("data1", outputsByKey.get("key1"));
+    assertEquals("finish_value", outputsByKey.get("finish_key"));
+
+    worker.stop();
+  }
+
+  @Test
+  public void 
testSingleKey_multiKeyBundleEnabled_finishBundleAttachesToBundleLevel()
+      throws Exception {
+    assumeTrue("Multi-key bundling is only supported in Streaming Engine", 
streamingEngine);
+
+    server.clearCommitsReceived();
+    StreamingDataflowWorker worker =
+        makeMultiKeyEnabledWorker(
+            new ConfigurableMultiKeyDoFn(
+                ImmutableMap.of("key1", ImmutableList.of("data1")),
+                ImmutableList.of(KV.of("finish_key", "finish_value"))));
+    worker.start();
+
+    String input =
+        "work {"
+            + "  computation_id: \""
+            + DEFAULT_COMPUTATION_ID
+            + "\""
+            + "  input_data_watermark: 0"
+            + "  work {"
+            + "    key: \"key1\""
+            + "    sharding_key: 1"
+            + "    work_token: 1"
+            + "    cache_token: 2"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data1\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "}";
+    Windmill.GetWorkResponse workResponse =
+        buildInput(
+            input,
+            CoderUtils.encodeToByteArray(
+                CollectionCoder.of(IntervalWindow.getCoder()),
+                Collections.singletonList(DEFAULT_WINDOW)));
+
+    
server.whenGetDataCalled().answerByDefault(StreamingDataflowWorkerTest::emptyDataResponder);
+    server.whenGetWorkCalled().thenReturn(workResponse);
+
+    Map<Long, Windmill.WorkItemCommitRequest> result = 
server.waitForAndGetCommits(1);
+    assertEquals(1, result.size());
+
+    assertTrue(result.containsKey(1L));
+    Windmill.WorkItemCommitRequest commit1 = result.get(1L);
+    assertEquals("key1", commit1.getKey().toStringUtf8());
+    // In multi-key mode, key1 only contains its own process outputs
+    assertEquals(1, commit1.getOutputMessagesCount());
+    Windmill.OutputMessageBundle outputBundle1 = commit1.getOutputMessages(0);
+    assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
outputBundle1.getDestinationStreamId());
+    assertEquals(1, outputBundle1.getBundlesCount());
+    assertEquals("key1", outputBundle1.getBundles(0).getKey().toStringUtf8());
+    assertEquals("data1", 
outputBundle1.getBundles(0).getMessages(0).getData().toStringUtf8());
+
+    // finishBundle outputs are flushed to the bundle level of 
MultiKeyWorkItemCommitRequest
+    List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+        server.getMultiKeyCommitsReceived();
+    assertEquals(1, multiKeyCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = 
multiKeyCommits.get(0);
+    assertEquals(1, multiKeyCommit.getOutputMessagesCount());
+    Windmill.OutputMessageBundle bundleLevel = 
multiKeyCommit.getOutputMessages(0);
+    assertEquals(DEFAULT_DESTINATION_STREAM_ID, 
bundleLevel.getDestinationStreamId());
+    assertEquals(1, bundleLevel.getBundlesCount());
+    assertEquals("finish_key", 
bundleLevel.getBundles(0).getKey().toStringUtf8());
+    assertEquals("finish_value", 
bundleLevel.getBundles(0).getMessages(0).getData().toStringUtf8());
+
+    worker.stop();
+  }
+
+  @Test
+  public void testMultiKey_threeKeys_withIntermediateEmptyKey() throws 
Exception {
+    assumeTrue("Multi-key bundling is only supported in Streaming Engine", 
streamingEngine);
+
+    server.clearCommitsReceived();
+    StreamingDataflowWorker worker =
+        makeMultiKeyEnabledWorker(
+            new ConfigurableMultiKeyDoFn(
+                ImmutableMap.of(
+                    "key1", ImmutableList.of("data1"),
+                    "key3", ImmutableList.of("data3")),
+                ImmutableList.of(KV.of("finish_key", "finish_value"))));
+    worker.start();
+
+    String batchInputText =
+        "work {"
+            + "  computation_id: \""
+            + DEFAULT_COMPUTATION_ID
+            + "\""
+            + "  input_data_watermark: 0"
+            + "  work {"
+            + "    key: \"key1\""
+            + "    sharding_key: 1"
+            + "    work_token: 1"
+            + "    cache_token: 2"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data1\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "  work {"
+            + "    key: \"key2\""
+            + "    sharding_key: 2"
+            + "    work_token: 2"
+            + "    cache_token: 3"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data2\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "  work {"
+            + "    key: \"key3\""
+            + "    sharding_key: 3"
+            + "    work_token: 3"
+            + "    cache_token: 4"
+            + "    key_group { high: 0 low: 1 }"
+            + "    message_bundles {"
+            + "      source_computation_id: \""
+            + DEFAULT_SOURCE_COMPUTATION_ID
+            + "\""
+            + "      messages {"
+            + "        timestamp: 0"
+            + "        data: \"data3\""
+            + "      }"
+            + "    }"
+            + "  }"
+            + "}";
+    Windmill.GetWorkResponse batchInput =
+        buildInput(
+            batchInputText,
+            CoderUtils.encodeToByteArray(
+                CollectionCoder.of(IntervalWindow.getCoder()),
+                Collections.singletonList(DEFAULT_WINDOW)));
+
+    
server.whenGetDataCalled().answerByDefault(StreamingDataflowWorkerTest::emptyDataResponder);
+    server.whenGetWorkCalled().thenReturn(batchInput);
+
+    Map<Long, Windmill.WorkItemCommitRequest> result = 
server.waitForAndGetCommits(3);
+    assertEquals(3, result.size());
+
+    // Verify Key 1 commit
+    assertTrue(result.containsKey(1L));
+    Windmill.WorkItemCommitRequest commit1 = result.get(1L);
+    assertEquals("key1", commit1.getKey().toStringUtf8());
+    assertEquals(1, commit1.getOutputMessagesCount());
+    assertEquals(
+        "data1",
+        
commit1.getOutputMessages(0).getBundles(0).getMessages(0).getData().toStringUtf8());
+
+    // Verify Key 2 commit (empty outputs)
+    assertTrue(result.containsKey(2L));
+    Windmill.WorkItemCommitRequest commit2 = result.get(2L);
+    assertEquals("key2", commit2.getKey().toStringUtf8());
+    assertEquals(0, commit2.getOutputMessagesCount());
+
+    // Verify Key 3 commit
+    assertTrue(result.containsKey(3L));
+    Windmill.WorkItemCommitRequest commit3 = result.get(3L);
+    assertEquals("key3", commit3.getKey().toStringUtf8());
+    assertEquals(1, commit3.getOutputMessagesCount());
+    assertEquals(
+        "data3",
+        
commit3.getOutputMessages(0).getBundles(0).getMessages(0).getData().toStringUtf8());
+
+    // Verify MultiKey commit: should contain all 3 requests and finishBundle 
outputs at bundle
+    // level
+    List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+        server.getMultiKeyCommitsReceived();
+    assertEquals(1, multiKeyCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = 
multiKeyCommits.get(0);
+    assertEquals(3, multiKeyCommit.getRequestsCount());
+    assertEquals(1, multiKeyCommit.getOutputMessagesCount());
+    Windmill.OutputMessageBundle bundleLevel = 
multiKeyCommit.getOutputMessages(0);
+    assertEquals("finish_key", 
bundleLevel.getBundles(0).getKey().toStringUtf8());
+    assertEquals("finish_value", 
bundleLevel.getBundles(0).getMessages(0).getData().toStringUtf8());
+
+    worker.stop();
+  }
+
   private StreamingDataflowWorker makeMultiKeyEnabledWorker() {
     return makeMultiKeyEnabledWorker(new WorkDoFn());
   }
@@ -5988,6 +6426,35 @@ public class StreamingDataflowWorkerTest {
     }
   }
 
+  static class ConfigurableMultiKeyDoFn extends DoFn<KV<String, String>, 
KV<String, String>> {
+    private final Map<String, List<String>> processOutputs;
+    private final List<KV<String, String>> finishBundleOutputs;
+
+    ConfigurableMultiKeyDoFn(
+        Map<String, List<String>> processOutputs, List<KV<String, String>> 
finishBundleOutputs) {
+      this.processOutputs = processOutputs;
+      this.finishBundleOutputs = finishBundleOutputs;
+    }
+
+    @ProcessElement
+    public void processElement(ProcessContext c) {
+      String key = c.element().getKey();
+      List<String> outputs = processOutputs.get(key);
+      if (outputs != null) {
+        for (String output : outputs) {
+          c.output(KV.of(key, output));
+        }
+      }
+    }
+
+    @FinishBundle
+    public void finishBundle(FinishBundleContext c) {
+      for (KV<String, String> output : finishBundleOutputs) {
+        c.output(output, new Instant(0), DEFAULT_WINDOW);
+      }
+    }
+  }
+
   @AutoValue
   abstract static class StreamingDataflowWorkerTestParams {
 
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java
index 043d02055cb..0251b075bec 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java
@@ -78,6 +78,7 @@ import 
org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDa
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache;
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV1;
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillTagEncodingV2;
+import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.ExecuteWorkResult;
 import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.StreamingEngineFailureTracker;
 import 
org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender;
 import org.apache.beam.sdk.Pipeline;
@@ -243,10 +244,11 @@ public class StreamingModeExecutionContextTest {
             TimeDomain.EVENT_TIME,
             CausedByDrain.NORMAL));
     executionContext.finishKey();
-    executionContext.flushState();
+    ExecuteWorkResult result = executionContext.flushStateAndReset();
 
-    Windmill.WorkItemCommitRequest.Builder outputBuilder = 
executionContext.getOutputBuilder();
-    Windmill.Timer timer = outputBuilder.buildPartial().getOutputTimers(0);
+    assertEquals(1, result.workItemCommits().size());
+    Windmill.WorkItemCommitRequest commitRequest = 
result.workItemCommits().get(0);
+    Windmill.Timer timer = commitRequest.getOutputTimers(0);
     assertThat(timer.getTag().toStringUtf8(), equalTo("/skey+0:5000"));
     assertThat(timer.getTimestamp(), 
equalTo(TimeUnit.MILLISECONDS.toMicros(5000)));
     assertThat(timer.getType(), equalTo(Windmill.Timer.Type.WATERMARK));
@@ -497,9 +499,10 @@ public class StreamingModeExecutionContextTest {
 
     stepContext.setBacklogBytes(1234.0);
     executionContext.finishKey();
-    executionContext.flushState();
+    ExecuteWorkResult result = executionContext.flushStateAndReset();
 
-    assertEquals(1234, 
executionContext.getOutputBuilder().getSourceBacklogBytes());
+    assertEquals(1, result.workItemCommits().size());
+    assertEquals(1234, 
result.workItemCommits().get(0).getSourceBacklogBytes());
   }
 
   @Test
@@ -866,7 +869,7 @@ public class StreamingModeExecutionContextTest {
     StateInternals stateInternals = stepContext.stateInternals();
 
     executionContext.finishKey();
-    executionContext.flushState();
+    executionContext.flushStateAndReset();
 
     // Verify timerInternals is poisoned
     try {
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java
index e69d8b5caa9..b72a218f1d0 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/WorkerCustomSourcesTest.java
@@ -105,6 +105,7 @@ import 
org.apache.beam.runners.dataflow.worker.util.common.worker.WorkExecutor;
 import org.apache.beam.runners.dataflow.worker.windmill.Windmill;
 import 
org.apache.beam.runners.dataflow.worker.windmill.client.getdata.FakeGetDataClient;
 import 
org.apache.beam.runners.dataflow.worker.windmill.state.WindmillStateCache;
+import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.ExecuteWorkResult;
 import 
org.apache.beam.runners.dataflow.worker.windmill.work.processing.failures.FailureTracker;
 import 
org.apache.beam.runners.dataflow.worker.windmill.work.refresh.HeartbeatSender;
 import org.apache.beam.sdk.Pipeline;
@@ -660,18 +661,19 @@ public class WorkerCustomSourcesTest {
     ByteString state = ByteString.EMPTY;
     for (int i = 0; i < 10 * maxElements;
     /* Incremented in inner loop */ ) {
+      Windmill.WorkItem workItem =
+          Windmill.WorkItem.newBuilder()
+              .setKey(ByteString.copyFromUtf8("0000000000000001")) // key is 
zero-padded index.
+              .setWorkToken(i) // Must be increasing across activations for 
cache to be used.
+              .setCacheToken(1)
+              .setSourceState(
+                  Windmill.SourceState.newBuilder().setState(state).build()) 
// Source state.
+              .build();
       // Initialize streaming context with state from previous iteration.
       startContext(
           context,
           createMockWork(
-              Windmill.WorkItem.newBuilder()
-                  .setKey(ByteString.copyFromUtf8("0000000000000001")) // key 
is zero-padded index.
-                  .setWorkToken(i) // Must be increasing across activations 
for cache to be used.
-                  .setCacheToken(1)
-                  .setSourceState(
-                      
Windmill.SourceState.newBuilder().setState(state).build()) // Source state.
-                  .build(),
-              Watermarks.builder().setInputDataWatermark(new 
Instant(0)).build()));
+              workItem, Watermarks.builder().setInputDataWatermark(new 
Instant(0)).build()));
 
       @SuppressWarnings({"unchecked", "rawtypes"})
       NativeReader<WindowedValue<ValueWithRecordId<KV<Integer, Integer>>>> 
reader =
@@ -706,21 +708,21 @@ public class WorkerCustomSourcesTest {
           numReadOnThisIteration, 
lessThanOrEqualTo(debugOptions.getUnboundedReaderMaxElements()));
 
       // Extract and verify state modifications.
-      context.flushState();
-      state = context.getOutputBuilder().getSourceStateUpdates().getState();
+      context.finishKey();
+      WindmillComputationKey computationKey = context.getComputationKey();
+      ExecuteWorkResult result = context.flushStateAndReset();
+      assertEquals(1, result.workItemCommits().size());
+      Windmill.WorkItemCommitRequest commitRequest = 
result.workItemCommits().get(0);
+      state = commitRequest.getSourceStateUpdates().getState();
       // CountingSource's watermark is the last record + 1.  i is now one past 
the last record,
       // so the expected watermark is i millis.
-      assertEquals(
-          TimeUnit.MILLISECONDS.toMicros(i), 
context.getOutputBuilder().getSourceWatermark());
-      assertEquals(
-          1, 
context.getOutputBuilder().getSourceStateUpdates().getFinalizeIdsList().size());
+      assertEquals(TimeUnit.MILLISECONDS.toMicros(i), 
commitRequest.getSourceWatermark());
+      assertEquals(1, 
commitRequest.getSourceStateUpdates().getFinalizeIdsList().size());
 
       assertNotNull(
           readerCache.acquireReader(
-              context.getComputationKey(),
-              context.getWorkItem().getCacheToken(),
-              context.getWorkToken() + 1));
-      assertEquals(7L, context.getBacklogBytes());
+              computationKey, workItem.getCacheToken(), 
workItem.getWorkToken() + 1));
+      assertEquals(7L, commitRequest.getSourceBacklogBytes());
     }
   }
 
diff --git 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperationTest.java
 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperationTest.java
index 2a83cc56cc0..95c7863d79f 100644
--- 
a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperationTest.java
+++ 
b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/util/common/worker/WriteOperationTest.java
@@ -21,9 +21,11 @@ import static 
org.apache.beam.runners.dataflow.worker.counters.CounterName.named
 import static org.hamcrest.CoreMatchers.hasItems;
 import static org.hamcrest.MatcherAssert.assertThat;
 import static org.hamcrest.Matchers.equalTo;
+import static org.junit.Assert.assertThrows;
 import static org.junit.Assert.fail;
 import static org.mockito.ArgumentMatchers.anyBoolean;
 import static org.mockito.ArgumentMatchers.eq;
+import static org.mockito.Mockito.doThrow;
 import static org.mockito.Mockito.mock;
 import static org.mockito.Mockito.never;
 import static org.mockito.Mockito.verify;
@@ -182,4 +184,61 @@ public class WriteOperationTest {
     inOrder.verify(sinkWriter).close();
     inOrder.verify(finishCloseable).close();
   }
+
+  @Test
+  public void testFinishKey() throws Exception {
+    OperationContext mockContext = mock(OperationContext.class);
+    when(mockContext.counterFactory()).thenReturn(counterSet);
+    
when(mockContext.nameContext()).thenReturn(NameContextsForTests.nameContextForTest());
+    Closeable startCloseable = mock(Closeable.class);
+    Closeable processCloseable = mock(Closeable.class);
+    when(mockContext.enterStart()).thenReturn(startCloseable);
+    when(mockContext.enterProcess()).thenReturn(processCloseable);
+
+    Sink sink = mock(Sink.class);
+    Sink.SinkWriter sinkWriter = mock(Sink.SinkWriter.class);
+    when(sink.writer()).thenReturn(sinkWriter);
+
+    WriteOperation operation = WriteOperation.forTest(sink, mockContext);
+    operation.start();
+    operation.finishKey("key1");
+
+    verify(mockContext).enterProcess();
+    verify(sinkWriter).finishKey("key1");
+    verify(processCloseable).close();
+  }
+
+  @Test
+  public void testFinishKey_unstarted_throwsException() throws Exception {
+    Sink sink = mock(Sink.class);
+    WriteOperation operation = WriteOperation.forTest(sink, context);
+
+    assertThrows(AssertionError.class, () -> operation.finishKey("key1"));
+  }
+
+  @Test
+  public void testFinishKey_nullKey() throws Exception {
+    Sink sink = mock(Sink.class);
+    Sink.SinkWriter sinkWriter = mock(Sink.SinkWriter.class);
+    when(sink.writer()).thenReturn(sinkWriter);
+
+    WriteOperation operation = WriteOperation.forTest(sink, context);
+    operation.start();
+    operation.finishKey(null);
+
+    verify(sinkWriter).finishKey(null);
+  }
+
+  @Test
+  public void testFinishKey_exceptionPropagates() throws Exception {
+    Sink sink = mock(Sink.class);
+    Sink.SinkWriter sinkWriter = mock(Sink.SinkWriter.class);
+    when(sink.writer()).thenReturn(sinkWriter);
+    doThrow(new IOException("finishKey 
error")).when(sinkWriter).finishKey("key1");
+
+    WriteOperation operation = WriteOperation.forTest(sink, context);
+    operation.start();
+
+    assertThrows(IOException.class, () -> operation.finishKey("key1"));
+  }
 }

Reply via email to