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


##########
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:
   should we throw an exception for now?



##########
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:
   since we are using this for single keys also that likely don't have a group, 
should we avoid setting the key group nested field if 0?



##########
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:
   can we throw an exception if multikey bundles are enabled to make sure we 
don't lose data if we forget to address this?



##########
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:
   might be interesting to capture metrics for status pages on why we stopped 
batching: poll was empty, too many items, too much time, sink
   
   Would help debug if we're not observing the bundling we expect



##########
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:
   use List.of for consistency with above (and seems clearer)
   
   ditto for rest of file



##########
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:
   // oldWork is null when newWork is the first work for the bundle.



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


Review Comment:
   but that key A and key C can be retried locally.



##########
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:
   add a helper to setup params for a multi-key worker?



##########
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:
   do we want this as part of the server? seems to be generic empty data 
response



##########
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:
   if we do add support for the generic response to the server, we could also 
add the ability to set work tokens to fail get data requests.



##########
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:
   do we need this harness threads because otherwise we start new threads to 
process the same group in parallel? 
   
   If so I wonder if that something we might want to try to prevent in the 
future? If we have more threads than # of CPU, it is likely better to not 
parallelize a key group.



##########
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:
   can you add a comment here?
   // We defer recordGetWorkStreamLatencies() to be called during bundle 
processing
   // as these are constructed on the hot GetWork thread



##########
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:
   advance again?



##########
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:
   validate invalidwork



##########
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:
   how about unifying the loops by just failing work if !isRunning (or cached 
value from before loop) before the if in the next loop?



##########
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:
   assert no interactions on executor



##########
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:
   List.of



-- 
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