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