Skip to content

Commit abbaab8

Browse files
authored
Fix flaky GrpcDirectGetWorkStreamTest.testConsumedWorkItems_itemsSplitAcrossResponses (#36129)
1 parent 2beb75c commit abbaab8

1 file changed

Lines changed: 16 additions & 11 deletions

File tree

runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/GrpcDirectGetWorkStreamTest.java

Lines changed: 16 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -392,7 +392,9 @@ public void testConsumedWorkItems() throws InterruptedException {
392392

393393
@Test
394394
public void testConsumedWorkItems_itemsSplitAcrossResponses() throws InterruptedException {
395-
int expectedRequests = 3;
395+
// We send all the responses on the first request. We don't care if there are additional
396+
// requests.
397+
int expectedRequests = 1;
396398
CountDownLatch waitForRequests = new CountDownLatch(expectedRequests);
397399
TestGetWorkRequestObserver requestObserver = new TestGetWorkRequestObserver(waitForRequests);
398400
GetWorkStreamTestStub testStub = new GetWorkStreamTestStub(requestObserver);
@@ -426,9 +428,9 @@ public void testConsumedWorkItems_itemsSplitAcrossResponses() throws Interrupted
426428
Windmill.WorkItem workItem3 =
427429
Windmill.WorkItem.newBuilder()
428430
.setKey(ByteString.copyFromUtf8("somewhat_long_key3"))
429-
.setWorkToken(2L)
430-
.setShardingKey(2L)
431-
.setCacheToken(2L)
431+
.setWorkToken(3L)
432+
.setShardingKey(3L)
433+
.setCacheToken(3L)
432434
.build();
433435

434436
List<ByteString> chunks1 = new ArrayList<>();
@@ -444,12 +446,12 @@ public void testConsumedWorkItems_itemsSplitAcrossResponses() throws Interrupted
444446

445447
chunks3.add(workItem3.toByteString());
446448

449+
assertTrue(waitForRequests.await(5, TimeUnit.SECONDS));
450+
447451
testStub.injectResponse(createResponse(chunks1, bytes.size() - third));
448452
testStub.injectResponse(createResponse(chunks2, bytes.size() - 2 * third));
449453
testStub.injectResponse(createResponse(chunks3, 0));
450454

451-
assertTrue(waitForRequests.await(5, TimeUnit.SECONDS));
452-
453455
assertThat(scheduledWorkItems).containsExactly(workItem1, workItem2, workItem3);
454456
}
455457

@@ -458,6 +460,7 @@ private static class GetWorkStreamTestStub
458460

459461
private final TestGetWorkRequestObserver requestObserver;
460462
private @Nullable StreamObserver<Windmill.StreamingGetWorkResponseChunk> responseObserver;
463+
private final CountDownLatch waitForStream = new CountDownLatch(1);
461464

462465
private GetWorkStreamTestStub(TestGetWorkRequestObserver requestObserver) {
463466
this.requestObserver = requestObserver;
@@ -466,15 +469,17 @@ private GetWorkStreamTestStub(TestGetWorkRequestObserver requestObserver) {
466469
@Override
467470
public StreamObserver<Windmill.StreamingGetWorkRequest> getWorkStream(
468471
StreamObserver<Windmill.StreamingGetWorkResponseChunk> responseObserver) {
469-
if (this.responseObserver == null) {
470-
this.responseObserver = responseObserver;
471-
requestObserver.responseObserver = this.responseObserver;
472-
}
472+
assertThat(this.responseObserver).isNull();
473+
this.responseObserver = responseObserver;
474+
requestObserver.responseObserver = this.responseObserver;
475+
waitForStream.countDown();
473476

474477
return requestObserver;
475478
}
476479

477-
private void injectResponse(Windmill.StreamingGetWorkResponseChunk responseChunk) {
480+
private void injectResponse(Windmill.StreamingGetWorkResponseChunk responseChunk)
481+
throws InterruptedException {
482+
waitForStream.await();
478483
checkNotNull(responseObserver).onNext(responseChunk);
479484
}
480485
}

0 commit comments

Comments
 (0)