arunpandianp commented on code in PR #38919:
URL: https://github.com/apache/beam/pull/38919#discussion_r3673964455


##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/streaming/Work.java:
##########
@@ -118,6 +120,7 @@ private Work(
             + Long.toHexString(workItem.getWorkToken());
     this.currentState = TimedState.initialState(startTime);
     this.isFailed = false;
+    this.getWorkStreamLatencies = getWorkStreamLatencies;

Review Comment:
   Thanks, done.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java:
##########
@@ -398,9 +386,52 @@ private void commitWorkBatch(
       ComputationState computationState,
       List<Work> workBatch,
       List<Windmill.WorkItemCommitRequest> workItemCommits) {
-    checkState(workBatch.size() == 1, "Expected single-key work batch, got: " 
+ workBatch.size());
-    checkState(workBatch.size() == workItemCommits.size());
-    commitSingleKeyWork(computationState, workBatch.get(0), 
workItemCommits.get(0));
+    if (workBatch.isEmpty()) {
+      return;
+    }
+    if (workBatch.size() > 1 || multiKeyBundleOptions.multiKeyBundleEnabled()) 
{
+      commitMultiKeyWorkBatch(computationState, workBatch, workItemCommits);
+    } else {
+      commitSingleKeyWork(computationState, workBatch.get(0), 
workItemCommits.get(0));
+    }
+  }
+
+  private void commitMultiKeyWorkBatch(
+      ComputationState computationState,
+      List<Work> workBatch,
+      List<Windmill.WorkItemCommitRequest> workItemCommits) {
+    Preconditions.checkState(!workBatch.isEmpty());
+    Preconditions.checkState(workBatch.size() == workItemCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest.Builder multiKeyBuilder =
+        Windmill.MultiKeyWorkItemCommitRequest.newBuilder();
+
+    Work primaryWork = workBatch.get(0);
+    Work.KeyGroup keyGroup = primaryWork.getKeyGroup();

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitter.java:
##########
@@ -145,11 +145,41 @@ private void drainCommitQueue() {
   }
 
   private void failQueuedCommit(Commit commit) {
+    if (!isRunning.get()) {
+      // Shutting down, fail everything unconditionally to prevent infinite 
loops
+      for (Work w : commit.workBatch()) {
+        w.setFailed();

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java:
##########
@@ -149,69 +135,120 @@ public void logAndProcessFailure_doesNotRetryOOM() {
     assertThrows(
         OutOfMemoryError.class,
         () ->
-            workFailureProcessor.logAndProcessFailure(
-                DEFAULT_COMPUTATION_ID, work, new OutOfMemoryError(), 
invalidWork::add));
+            workFailureProcessor.logAndProcessFailureBatch(
+                DEFAULT_COMPUTATION_ID,
+                Arrays.asList(work),
+                new OutOfMemoryError(),
+                invalidWork::add));
 
     assertThat(executedWork).isEmpty();

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java:
##########
@@ -1252,40 +1253,368 @@ public void 
testNumberOfWorkerHarnessThreadsIsHonored() throws Exception {
   }
 
   @Test
-  public void testKeyTokenInvalidException() throws Exception {
-    if (streamingEngine) {
-      // TODO: This test needs to be adapted to work with streamingEngine=true.
+  public void testMultiKeyCommit_success() throws Exception {
+    if (!streamingEngine) {
       return;
     }
     KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
 
     List<ParallelInstruction> instructions =
         Arrays.asList(
             makeSourceInstruction(kvCoder),
-            makeDoFnInstruction(new KeyTokenInvalidFn(), 0, kvCoder),
+            makeDoFnInstruction(new WorkDoFn(), 0, kvCoder),
             makeSinkInstruction(kvCoder, 1));
 
+    StreamingDataflowWorker worker =
+        makeWorker(

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java:
##########
@@ -1252,40 +1253,368 @@ public void 
testNumberOfWorkerHarnessThreadsIsHonored() throws Exception {
   }
 
   @Test
-  public void testKeyTokenInvalidException() throws Exception {
-    if (streamingEngine) {
-      // TODO: This test needs to be adapted to work with streamingEngine=true.
+  public void testMultiKeyCommit_success() throws Exception {
+    if (!streamingEngine) {
       return;
     }
     KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
 
     List<ParallelInstruction> instructions =
         Arrays.asList(
             makeSourceInstruction(kvCoder),
-            makeDoFnInstruction(new KeyTokenInvalidFn(), 0, kvCoder),
+            makeDoFnInstruction(new WorkDoFn(), 0, kvCoder),
             makeSinkInstruction(kvCoder, 1));
 
+    StreamingDataflowWorker worker =
+        makeWorker(
+            defaultWorkerParams(
+                    
"--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000",
+                    "--numberOfWorkerHarnessThreads=1")

Review Comment:
   Yes, there is room to improve the batching. Planning to tackle it separately.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java:
##########
@@ -187,7 +188,7 @@ public class StreamingModeExecutionContext
 
   // Key switch listener to delegate MDC logging context and thread name 
updates
   public interface KeyTransitionListener {
-    void onKeyTransition(Work oldWork, Work newWork);
+    void onKeyTransition(@Nullable Work oldWork, Work newWork);

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java:
##########
@@ -508,6 +517,274 @@ public void testStart_internalKeyDecoding() throws 
Exception {
     assertEquals("decodedKey", executionContext.getKey());
   }
 
+  @Test
+  public void testAdvance_success() throws Exception {
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    Windmill.WorkItem workItem2 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key2"))
+            .setWorkToken(2L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work2 =
+        createMockWork(
+            workItem2, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+    ExecutableWork executableWork2 = ExecutableWork.create(work2, (w, h) -> 
{});
+
+    org.mockito.Mockito.when(
+            mockExecutor.pollWork(
+                org.mockito.Mockito.eq(COMPUTATION_ID),
+                org.mockito.Mockito.eq(work1.getKeyGroup()),
+                org.mockito.Mockito.eq(mockHandle)))
+        .thenReturn(executableWork2);
+
+    executionContext.start(
+        work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, 
newWork) -> {});
+
+    assertTrue(executionContext.advance());
+    assertEquals("key2", executionContext.getSerializedKey().toStringUtf8());
+  }
+
+  @Test
+  public void testAdvance_noMoreWork() throws Exception {
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    org.mockito.Mockito.when(
+            mockExecutor.pollWork(
+                org.mockito.Mockito.eq(COMPUTATION_ID),
+                org.mockito.Mockito.eq(work1.getKeyGroup()),
+                org.mockito.Mockito.eq(mockHandle)))
+        .thenReturn(null);
+
+    executionContext.start(
+        work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, 
newWork) -> {});
+
+    assertFalse(executionContext.advance());
+  }
+
+  @Test
+  public void testAdvance_respectsMaxBatchSize() throws Exception {
+    DataflowWorkerHarnessOptions optionsWithBatchSize =
+        PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class);
+    optionsWithBatchSize
+        .as(ExperimentalOptions.class)
+        .setExperiments(Arrays.asList("windmill_max_key_group_batch_size=1"));
+    StreamingModeExecutionContext context =
+        createExecutionContext(optionsWithBatchSize, globalConfigHandle);
+
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    context.start(work1, workExecutor, mockExecutor, mockHandle, null, 
(oldWork, newWork) -> {});
+
+    assertFalse(context.advance());
+    org.mockito.Mockito.verifyNoInteractions(mockExecutor);
+  }
+
+  @Test
+  public void testAdvance_respectsMaxBatchTime() throws Exception {
+    DataflowWorkerHarnessOptions optionsWithBatchTime =
+        PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class);
+    optionsWithBatchTime
+        .as(ExperimentalOptions.class)
+        
.setExperiments(Arrays.asList("windmill_max_key_group_batch_time_ms=0"));
+    StreamingModeExecutionContext context =
+        createExecutionContext(optionsWithBatchTime, globalConfigHandle);
+
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    context.start(work1, workExecutor, mockExecutor, mockHandle, null, 
(oldWork, newWork) -> {});
+
+    assertFalse(context.advance());
+    org.mockito.Mockito.verifyNoInteractions(mockExecutor);
+  }
+
+  @Test
+  public void testAdvance_workFailed() throws Exception {
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    executionContext.start(
+        work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, 
newWork) -> {});
+
+    work1.setFailed();
+
+    assertThrows(WorkItemCancelledException.class, () -> 
executionContext.advance());

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/StreamingWorkScheduler.java:
##########
@@ -398,9 +386,52 @@ private void commitWorkBatch(
       ComputationState computationState,
       List<Work> workBatch,
       List<Windmill.WorkItemCommitRequest> workItemCommits) {
-    checkState(workBatch.size() == 1, "Expected single-key work batch, got: " 
+ workBatch.size());
-    checkState(workBatch.size() == workItemCommits.size());
-    commitSingleKeyWork(computationState, workBatch.get(0), 
workItemCommits.get(0));
+    if (workBatch.isEmpty()) {
+      return;
+    }
+    if (workBatch.size() > 1 || multiKeyBundleOptions.multiKeyBundleEnabled()) 
{
+      commitMultiKeyWorkBatch(computationState, workBatch, workItemCommits);
+    } else {
+      commitSingleKeyWork(computationState, workBatch.get(0), 
workItemCommits.get(0));
+    }
+  }
+
+  private void commitMultiKeyWorkBatch(
+      ComputationState computationState,
+      List<Work> workBatch,
+      List<Windmill.WorkItemCommitRequest> workItemCommits) {
+    Preconditions.checkState(!workBatch.isEmpty());
+    Preconditions.checkState(workBatch.size() == workItemCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest.Builder multiKeyBuilder =
+        Windmill.MultiKeyWorkItemCommitRequest.newBuilder();
+
+    Work primaryWork = workBatch.get(0);
+    Work.KeyGroup keyGroup = primaryWork.getKeyGroup();
+    multiKeyBuilder.setKeyGroup(
+        
Windmill.Uint128Proto.newBuilder().setHigh(keyGroup.high()).setLow(keyGroup.low()).build());
+
+    for (int i = 0; i < workBatch.size(); i++) {
+      // TODO: Retry on commit truncations

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java:
##########
@@ -739,12 +752,51 @@ private final long computeSourceBytesProcessed(String 
sourceBytesCounterName) {
         .orElse(0L);
   }
 
-  public boolean advance() {
-    // TODO: get more work from workQueueExecutor and merge into the bundle 
here
+  public boolean advance() throws CoderException {
+    if (!multiKeyBundleOptions.multiKeyBundleEnabled()) {
+      return false;
+    }
+
+    Work activeWork = checkStateNotNull(work);
+    BoundedQueueExecutor executor = checkStateNotNull(workQueueExecutor);
+    BoundedQueueExecutorWorkHandle handle = checkStateNotNull(budgetHandle);
+
+    if (workIsFailed()) {
+      throw new 
WorkItemCancelledException(activeWork.getWorkItem().getShardingKey());
+    }
+
+    if (activeWork.getKeyGroup().equals(Work.KeyGroup.DEFAULT) || 
shouldStopBatching()) {
+      return false;
+    }
+
+    @Nullable
+    ExecutableWork additionalWork =
+        executor.pollWork(computationId, activeWork.getKeyGroup(), handle);
+    if (additionalWork != null) {
+      flushStateInternal();
+      Work newWork = additionalWork.work();
+      ++workItemsPolled;
+      checkStateNotNull(keyTransitionListener).onKeyTransition(activeWork, 
newWork);
+      startForNewKey(newWork);
+      return true;
+    }
+
     return false;
   }
 
-  private void startForNewKey(Work newWork, WindmillStateReader reader) throws 
CoderException {
+  private boolean shouldStopBatching() {

Review Comment:
   +1 I'll take it as a follow up item. Don't want to block the PR on it.



##########
runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContext.java:
##########
@@ -714,6 +725,8 @@ private void validateCommitRequestSize() {
         buildWorkItemTruncationRequestBuilder(currentWork, 
estimatedCommitSize);
     currentBuilder.clear();
     currentBuilder.mergeFrom(truncationBuilder.build());
+
+    // TODO: throw and retry when truncation is not on a single key bundle.

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/work/processing/failures/WorkFailureProcessorTest.java:
##########
@@ -149,69 +135,120 @@ public void logAndProcessFailure_doesNotRetryOOM() {
     assertThrows(
         OutOfMemoryError.class,
         () ->
-            workFailureProcessor.logAndProcessFailure(
-                DEFAULT_COMPUTATION_ID, work, new OutOfMemoryError(), 
invalidWork::add));
+            workFailureProcessor.logAndProcessFailureBatch(
+                DEFAULT_COMPUTATION_ID,
+                Arrays.asList(work),

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/commits/StreamingEngineWorkCommitterTest.java:
##########


Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java:
##########
@@ -1252,40 +1253,368 @@ public void 
testNumberOfWorkerHarnessThreadsIsHonored() throws Exception {
   }
 
   @Test
-  public void testKeyTokenInvalidException() throws Exception {
-    if (streamingEngine) {
-      // TODO: This test needs to be adapted to work with streamingEngine=true.
+  public void testMultiKeyCommit_success() throws Exception {
+    if (!streamingEngine) {
       return;
     }
     KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
 
     List<ParallelInstruction> instructions =
         Arrays.asList(
             makeSourceInstruction(kvCoder),
-            makeDoFnInstruction(new KeyTokenInvalidFn(), 0, kvCoder),
+            makeDoFnInstruction(new WorkDoFn(), 0, kvCoder),
             makeSinkInstruction(kvCoder, 1));
 
+    StreamingDataflowWorker worker =
+        makeWorker(
+            defaultWorkerParams(
+                    
"--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000",
+                    "--numberOfWorkerHarnessThreads=1")
+                .setLocalRetryTimeoutMs(100)
+                .setInstructions(instructions)
+                .build());
+    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
-        .whenGetWorkCalled()
-        .thenReturn(makeInput(0, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY));
+        .whenGetDataCalled()
+        .answerByDefault(
+            request -> {
+              Windmill.GetDataResponse.Builder builder = 
Windmill.GetDataResponse.newBuilder();
+              for (ComputationGetDataRequest compRequest : 
request.getRequestsList()) {
+                ComputationGetDataResponse.Builder compBuilder =
+                    
builder.addDataBuilder().setComputationId(compRequest.getComputationId());
+                for (KeyedGetDataRequest keyRequest : 
compRequest.getRequestsList()) {
+                  KeyedGetDataResponse.Builder keyBuilder =
+                      compBuilder
+                          .addDataBuilder()
+                          .setKey(keyRequest.getKey())
+                          .setShardingKey(keyRequest.getShardingKey());
+                  keyBuilder.addAllValues(keyRequest.getValuesToFetchList());
+                  keyBuilder.addAllBags(keyRequest.getBagsToFetchList());
+                  
keyBuilder.addAllWatermarkHolds(keyRequest.getWatermarkHoldsToFetchList());
+                }
+              }
+              return builder.build();
+            });
+
+    server.whenGetWorkCalled().thenReturn(batchInput);
+
+    Map<Long, Windmill.WorkItemCommitRequest> result = 
server.waitForAndGetCommits(3);
+
+    assertEquals(3, result.size());
+
+    List<Windmill.MultiKeyWorkItemCommitRequest> multiKeyCommits =
+        server.getMultiKeyCommitsReceived();
+    assertEquals(1, multiKeyCommits.size());
+    Windmill.MultiKeyWorkItemCommitRequest multiKeyCommit = 
multiKeyCommits.get(0);
+    assertEquals(3, multiKeyCommit.getRequestsCount());
+    assertEquals(1, multiKeyCommit.getRequests(0).getWorkToken());
+    assertEquals(2, multiKeyCommit.getRequests(1).getWorkToken());
+    assertEquals(3, multiKeyCommit.getRequests(2).getWorkToken());
+
+    worker.stop();
+  }
+
+  @Test
+  public void testMultiKeyCommit_elementFailure() throws Exception {
+    if (!streamingEngine) {
+      return;
+    }
+    KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
+
+    List<ParallelInstruction> instructions =
+        Arrays.asList(
+            makeSourceInstruction(kvCoder),
+            makeDoFnInstruction(new WorkDoFn(), 0, kvCoder),
+            makeSinkInstruction(kvCoder, 1));
 
     StreamingDataflowWorker worker =
-        
makeWorker(defaultWorkerParams().setInstructions(instructions).publishCounters().build());
+        makeWorker(
+            defaultWorkerParams(
+                    
"--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=5000",
+                    "--numberOfWorkerHarnessThreads=1")
+                .setLocalRetryTimeoutMs(100)
+                .setInstructions(instructions)
+                .build());
     worker.start();
 
-    server.waitForEmptyWorkQueue();
+    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
-        .whenGetWorkCalled()
-        .thenReturn(makeInput(1, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY));
+        .whenGetDataCalled()
+        .answerByDefault(
+            request -> {
+              Windmill.GetDataResponse.Builder builder = 
Windmill.GetDataResponse.newBuilder();
+              for (ComputationGetDataRequest compRequest : 
request.getRequestsList()) {
+                ComputationGetDataResponse.Builder compBuilder =
+                    
builder.addDataBuilder().setComputationId(compRequest.getComputationId());
+                for (KeyedGetDataRequest keyRequest : 
compRequest.getRequestsList()) {
+                  KeyedGetDataResponse.Builder keyBuilder =
+                      compBuilder
+                          .addDataBuilder()
+                          .setKey(keyRequest.getKey())
+                          .setShardingKey(keyRequest.getShardingKey());
+                  if (keyRequest.getWorkToken() == 2) {
+                    keyBuilder.setFailed(true);

Review Comment:
   Added StreamingDataflowWorkerTest::emptyDataResponderWithFailedWorkTokens 
that take the failed work tokens.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java:
##########
@@ -508,6 +517,274 @@ public void testStart_internalKeyDecoding() throws 
Exception {
     assertEquals("decodedKey", executionContext.getKey());
   }
 
+  @Test
+  public void testAdvance_success() throws Exception {
+    BoundedQueueExecutor mockExecutor = mock(BoundedQueueExecutor.class);
+    BoundedQueueExecutorWorkHandle mockHandle = 
mock(BoundedQueueExecutorWorkHandle.class);
+
+    Windmill.Uint128Proto keyGroup =
+        Windmill.Uint128Proto.newBuilder().setHigh(1).setLow(2).build();
+    Windmill.WorkItem workItem1 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key1"))
+            .setWorkToken(1L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work1 =
+        createMockWork(
+            workItem1, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+
+    Windmill.WorkItem workItem2 =
+        Windmill.WorkItem.newBuilder()
+            .setKey(ByteString.copyFromUtf8("key2"))
+            .setWorkToken(2L)
+            .setKeyGroup(keyGroup)
+            .build();
+    Work work2 =
+        createMockWork(
+            workItem2, 
Watermarks.builder().setInputDataWatermark(Instant.EPOCH).build());
+    ExecutableWork executableWork2 = ExecutableWork.create(work2, (w, h) -> 
{});
+
+    org.mockito.Mockito.when(
+            mockExecutor.pollWork(
+                org.mockito.Mockito.eq(COMPUTATION_ID),
+                org.mockito.Mockito.eq(work1.getKeyGroup()),
+                org.mockito.Mockito.eq(mockHandle)))
+        .thenReturn(executableWork2);
+
+    executionContext.start(
+        work1, workExecutor, mockExecutor, mockHandle, null, (oldWork, 
newWork) -> {});
+
+    assertTrue(executionContext.advance());
+    assertEquals("key2", executionContext.getSerializedKey().toStringUtf8());

Review Comment:
   done.



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingDataflowWorkerTest.java:
##########
@@ -1252,40 +1253,368 @@ public void 
testNumberOfWorkerHarnessThreadsIsHonored() throws Exception {
   }
 
   @Test
-  public void testKeyTokenInvalidException() throws Exception {
-    if (streamingEngine) {
-      // TODO: This test needs to be adapted to work with streamingEngine=true.
+  public void testMultiKeyCommit_success() throws Exception {
+    if (!streamingEngine) {
       return;
     }
     KvCoder<String, String> kvCoder = KvCoder.of(StringUtf8Coder.of(), 
StringUtf8Coder.of());
 
     List<ParallelInstruction> instructions =
         Arrays.asList(
             makeSourceInstruction(kvCoder),
-            makeDoFnInstruction(new KeyTokenInvalidFn(), 0, kvCoder),
+            makeDoFnInstruction(new WorkDoFn(), 0, kvCoder),
             makeSinkInstruction(kvCoder, 1));
 
+    StreamingDataflowWorker worker =
+        makeWorker(
+            defaultWorkerParams(
+                    
"--experiments=unstable_enable_multi_key_bundle,windmill_max_key_group_batch_time_ms=50000",
+                    "--numberOfWorkerHarnessThreads=1")
+                .setLocalRetryTimeoutMs(100)
+                .setInstructions(instructions)
+                .build());
+    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
-        .whenGetWorkCalled()
-        .thenReturn(makeInput(0, 0, DEFAULT_KEY_STRING, DEFAULT_SHARDING_KEY));
+        .whenGetDataCalled()
+        .answerByDefault(

Review Comment:
   replaced it with the 
existing`StreamingDataflowWorkerTest::emptyDataResponder`. 



##########
runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/StreamingModeExecutionContextTest.java:
##########
@@ -148,15 +153,19 @@ private StreamingModeExecutionContext 
createExecutionContext(
         StreamingCounters.create(),
         mock(FailureTracker.class),
         "sourceBytesProcessCounterName",
+        MultiKeyBundleOptions.fromOptions(options),
         SideInputStateFetcherFactory.fromOptions(options));
   }
 
   @Before
   public void setUp() {
     MockitoAnnotations.initMocks(this);
     options = PipelineOptionsFactory.as(DataflowWorkerHarnessOptions.class);
+    options
+        .as(ExperimentalOptions.class)
+        .setExperiments(Arrays.asList("unstable_enable_multi_key_bundle"));

Review Comment:
   done.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to