diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserver.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserver.java index 1e197c877d68..b027a6cac7b0 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserver.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserver.java @@ -115,14 +115,24 @@ public void onNext(T t) throws StreamClosedException, WindmillStreamShutdownExce logger.debug("Stream was shutdown during send.", cancellationException); return; } + if (delegateStreamObserver == delegate) { + if (isCurrentStreamClosed) { + logger.debug("Stream is already closed when encountering error with send."); + return; + } + isCurrentStreamClosed = true; + } } + // Either this was the active observer the current observer that requires closing, or this was + // a previous + // observer which we attempt to close and ignore possible exceptions. try { delegate.onError(cancellationException); } catch (IllegalStateException onErrorException) { // The delegate above was already terminated via onError or onComplete. - // Fallthrough since this is possibly due to queued onNext() calls that are being made from - // previously blocked threads. + // Fallthrough since this is possibly due to queued onNext() calls that are being made + // from previously blocked threads. } catch (RuntimeException onErrorException) { logger.warn( "Encountered unexpected error {} when cancelling due to error.", @@ -134,14 +144,20 @@ public void onNext(T t) throws StreamClosedException, WindmillStreamShutdownExce public synchronized void onError(Throwable throwable) throws StreamClosedException, WindmillStreamShutdownException { - delegate().onError(throwable); - isCurrentStreamClosed = true; + try { + delegate().onError(throwable); + } finally { + isCurrentStreamClosed = true; + } } public synchronized void onCompleted() throws StreamClosedException, WindmillStreamShutdownException { - delegate().onCompleted(); - isCurrentStreamClosed = true; + try { + delegate().onCompleted(); + } finally { + isCurrentStreamClosed = true; + } } synchronized boolean isClosed() { diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/DirectStreamObserver.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/DirectStreamObserver.java index 173cbd26c4e7..bf060bd6acfe 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/DirectStreamObserver.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/DirectStreamObserver.java @@ -182,8 +182,8 @@ public void onError(Throwable t) { Preconditions.checkState(!isUserClosed); isUserClosed = true; if (!isOutboundObserverClosed) { - outboundObserver.onError(t); isOutboundObserverClosed = true; + outboundObserver.onError(t); } } } diff --git a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/StreamObserverCancelledException.java b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/StreamObserverCancelledException.java index 70fd3497a37f..5682d5085d2b 100644 --- a/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/StreamObserverCancelledException.java +++ b/runners/google-cloud-dataflow-java/worker/src/main/java/org/apache/beam/runners/dataflow/worker/windmill/client/grpc/observers/StreamObserverCancelledException.java @@ -21,15 +21,15 @@ @Internal public final class StreamObserverCancelledException extends RuntimeException { - StreamObserverCancelledException(Throwable cause) { + public StreamObserverCancelledException(Throwable cause) { super(cause); } - StreamObserverCancelledException(String message, Throwable cause) { + public StreamObserverCancelledException(String message, Throwable cause) { super(message, cause); } - StreamObserverCancelledException(String message) { + public StreamObserverCancelledException(String message) { super(message); } } diff --git a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserverTest.java b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserverTest.java index ef7a865748dd..69c54a50b574 100644 --- a/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserverTest.java +++ b/runners/google-cloud-dataflow-java/worker/src/test/java/org/apache/beam/runners/dataflow/worker/windmill/client/ResettableThrowingStreamObserverTest.java @@ -18,12 +18,15 @@ package org.apache.beam.runners.dataflow.worker.windmill.client; import static org.junit.Assert.assertThrows; +import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.eq; import static org.mockito.ArgumentMatchers.isA; +import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.spy; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.verifyNoInteractions; +import org.apache.beam.runners.dataflow.worker.windmill.client.grpc.observers.StreamObserverCancelledException; import org.apache.beam.runners.dataflow.worker.windmill.client.grpc.observers.TerminatingStreamObserver; import org.junit.Test; import org.junit.runner.RunWith; @@ -51,6 +54,53 @@ public void terminate(Throwable terminationException) {} }); } + @Test + public void testOnNext_simple() throws Exception { + ResettableThrowingStreamObserver observer = newStreamObserver(); + TerminatingStreamObserver spiedDelegate = newDelegate(); + observer.reset(spiedDelegate); + observer.onNext(1); + verify(spiedDelegate).onNext(eq(1)); + observer.onNext(2); + verify(spiedDelegate).onNext(eq(2)); + observer.onCompleted(); + verify(spiedDelegate).onCompleted(); + } + + @Test + public void testOnError_success() throws Exception { + ResettableThrowingStreamObserver observer = newStreamObserver(); + TerminatingStreamObserver spiedDelegate = newDelegate(); + observer.reset(spiedDelegate); + Throwable t = new RuntimeException("Test exception"); + observer.onError(t); + verify(spiedDelegate).onError(eq(t)); + + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, () -> observer.onNext(1)); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, observer::onCompleted); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, + () -> observer.onError(new RuntimeException("ignored"))); + } + + @Test + public void testOnCompleted_success() throws Exception { + ResettableThrowingStreamObserver observer = newStreamObserver(); + TerminatingStreamObserver spiedDelegate = newDelegate(); + observer.reset(spiedDelegate); + observer.onCompleted(); + verify(spiedDelegate).onCompleted(); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, () -> observer.onNext(1)); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, observer::onCompleted); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, + () -> observer.onError(new RuntimeException("ignored"))); + } + @Test public void testPoison_beforeDelegateSet() { ResettableThrowingStreamObserver observer = newStreamObserver(); @@ -97,9 +147,7 @@ public void testOnCompleted_afterPoisonedThrows() { } @Test - public void testReset_usesNewDelegate() - throws WindmillStreamShutdownException, - ResettableThrowingStreamObserver.StreamClosedException { + public void testReset_usesNewDelegate() throws Exception { ResettableThrowingStreamObserver observer = newStreamObserver(); TerminatingStreamObserver firstObserver = newDelegate(); observer.reset(firstObserver); @@ -113,6 +161,24 @@ public void testReset_usesNewDelegate() verify(secondObserver).onNext(eq(2)); } + @Test + public void testOnNext_streamCancelledException_closesStream() throws Exception { + ResettableThrowingStreamObserver observer = newStreamObserver(); + TerminatingStreamObserver spiedDelegate = newDelegate(); + StreamObserverCancelledException streamObserverCancelledException = + new StreamObserverCancelledException("Test error"); + doThrow(streamObserverCancelledException).when(spiedDelegate).onNext(any()); + observer.reset(spiedDelegate); + observer.onNext(1); + + verify(spiedDelegate).onError(eq(streamObserverCancelledException)); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, + () -> observer.onError(new Exception())); + assertThrows( + ResettableThrowingStreamObserver.StreamClosedException.class, observer::onCompleted); + } + private ResettableThrowingStreamObserver newStreamObserver() { return new ResettableThrowingStreamObserver<>(LoggerFactory.getLogger(getClass())); }