diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractReadContext.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractReadContext.java index c473d0719f49..901890b37d75 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractReadContext.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractReadContext.java @@ -1301,6 +1301,7 @@ CloseableIterator startStream( request.getLastStatement(), prefetchChunks, cancelQueryWhenClientIsClosed); + setStream(stream); if (streamListener != null) { stream.registerListener(streamListener); } @@ -1536,6 +1537,7 @@ CloseableIterator startStream( GrpcStreamIterator stream = new GrpcStreamIterator( lastStatement, prefetchChunks, cancelQueryWhenClientIsClosed); + setStream(stream); if (streamListener != null) { stream.registerListener(streamListener); } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractResultSet.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractResultSet.java index 0717cae74f25..9019cacab916 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractResultSet.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AbstractResultSet.java @@ -166,6 +166,14 @@ default boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMes /** it requests the initial prefetch chunks from gRPC stream */ default void requestPrefetchChunks() {} + + /** + * Returns true if data (a chunk, row, EOF, or error) is available to read immediately without + * blocking the calling thread on network I/O. + */ + default boolean isDataAvailable() { + return true; + } } static double valueProtoToFloat64(com.google.protobuf.Value proto) { diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AsyncResultSetImpl.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AsyncResultSetImpl.java index 3dd5724532b3..5feb058c2d64 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AsyncResultSetImpl.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/AsyncResultSetImpl.java @@ -20,7 +20,7 @@ import com.google.api.core.ApiFutures; import com.google.api.core.SettableApiFuture; import com.google.api.gax.core.ExecutorProvider; -import com.google.cloud.spanner.AbstractReadContext.ListenableAsyncResultSet; +import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Function; import com.google.common.base.Preconditions; import com.google.common.base.Supplier; @@ -35,21 +35,23 @@ import java.util.LinkedList; import java.util.List; import java.util.concurrent.BlockingDeque; -import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutionException; import java.util.concurrent.Executor; import java.util.concurrent.Future; import java.util.concurrent.LinkedBlockingDeque; import java.util.logging.Level; import java.util.logging.Logger; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; /** Default implementation for {@link AsyncResultSet}. */ +@NullMarked class AsyncResultSetImpl extends ForwardingStructReader - implements ListenableAsyncResultSet, AsyncResultSet.StreamMessageListener { + implements AbstractReadContext.ListenableAsyncResultSet, AsyncResultSet.StreamMessageListener { private static final Logger log = Logger.getLogger(AsyncResultSetImpl.class.getName()); /** State of an {@link AsyncResultSetImpl}. */ - private enum State { + enum State { INITIALIZED, STREAMING_INITIALIZED, /** SYNC indicates that the {@link ResultSet} is used in sync pattern. */ @@ -90,7 +92,7 @@ private enum State { private final ListeningScheduledExecutorService service; private final BlockingDeque buffer; - private Struct currentRow; + @Nullable private Struct currentRow; /** Supplies the underlying synchronous {@link ResultSet} that will be producing the rows. */ private final Supplier delegateResultSet; @@ -99,15 +101,15 @@ private enum State { * Any exception that occurs while executing the query and iterating over the result set will be * stored in this variable and propagated to the user through {@link #tryNext()}. */ - private volatile SpannerException executionException; + @Nullable private volatile SpannerException executionException; /** * Executor for callbacks. Regardless of the type of executor that is provided, the {@link * AsyncResultSetImpl} will ensure that at most 1 callback call will be active at any one time. */ - private Executor executor; + @Nullable private Executor executor; - private ReadyCallback callback; + @Nullable private ReadyCallback callback; /** * Listeners that will be called when the {@link AsyncResultSetImpl} has finished fetching all @@ -117,15 +119,24 @@ private enum State { private volatile State state = State.INITIALIZED; - /** This variable indicates that produce rows thread is initiated */ - private volatile boolean produceRowsInitiated; + /** Indicates whether a task is currently executing ProduceRowsRunnable. */ + private boolean producerRunning; + + /** Indicates whether a ProduceRowsRunnable task has been submitted and is waiting to run. */ + private boolean producerScheduled; + + /** Indicates whether produce rows has been initiated. */ + private boolean produceRowsInitiated; + + /** Indicates whether the result future and cleanup have been completed. */ + private boolean completed; /** * This variable indicates whether all the results from the underlying result set have been read. */ private volatile boolean finished; - private volatile SettableApiFuture result; + @Nullable private volatile SettableApiFuture result; /** * This variable indicates whether {@link #tryNext()} has returned {@link CursorState#DONE} or a @@ -133,23 +144,9 @@ private enum State { */ private volatile boolean cursorReturnedDoneOrException; - /** - * This variable is used to pause the producer when the {@link AsyncResultSet} is paused. The - * production of rows that are put into the buffer is only paused once the buffer is full. - */ - private volatile CountDownLatch pausedLatch = new CountDownLatch(1); - - /** - * This variable is used to pause the producer when the buffer is full and the consumer needs some - * time to catch up. - */ - private volatile CountDownLatch bufferConsumptionLatch = new CountDownLatch(0); - - /** - * This variable is used to pause the producer when all rows have been put into the buffer, but - * the consumer (the callback) has not yet received and processed all rows. - */ - private volatile CountDownLatch consumingLatch = new CountDownLatch(0); + private boolean callbackRunning; + private boolean callbackScheduled; + private boolean resumeRequestedWhileConsuming; AsyncResultSetImpl(ExecutorProvider executorProvider, ResultSet delegate, int bufferSize) { this(executorProvider, Suppliers.ofInstance(Preconditions.checkNotNull(delegate)), bufferSize); @@ -191,15 +188,26 @@ boolean isUsed() { */ @Override public void close() { + boolean shouldCloseDelegate = false; + boolean shouldShutdownService = false; synchronized (monitor) { if (this.closed) { return; } if (state == State.INITIALIZED || state == State.SYNC) { - delegateResultSet.get().close(); + shouldCloseDelegate = true; + if (executorProvider.shouldAutoClose()) { + shouldShutdownService = true; + } } this.closed = true; } + if (shouldCloseDelegate) { + closeDelegateResultSet(); + } + if (shouldShutdownService) { + service.shutdown(); + } } /** @@ -234,8 +242,6 @@ public CursorState tryNext() throws SpannerException { cursorReturnedDoneOrException = true; throw executionException; } - Preconditions.checkState( - this.callback != null, "tryNext may only be called after a callback has been set."); Preconditions.checkState( this.state == State.CONSUMING, "tryNext may only be called from a DataReady callback. Current state: " @@ -246,12 +252,11 @@ public CursorState tryNext() throws SpannerException { return CursorState.DONE; } } - if (!buffer.isEmpty()) { + Struct nextRow = buffer.poll(); + if (nextRow != null) { // Set the next row from the buffer as the current row of the StructReader. - replaceDelegate(currentRow = buffer.pop()); - synchronized (monitor) { - bufferConsumptionLatch.countDown(); - } + replaceDelegate(currentRow = nextRow); + scheduleProducerIfNecessary(); return CursorState.OK; } return CursorState.NOT_READY; @@ -272,8 +277,14 @@ private void closeDelegateResultSet() { private class CallbackRunnable implements Runnable { @Override public void run() { + synchronized (monitor) { + callbackScheduled = false; + callbackRunning = true; + } + boolean shouldScheduleProducer = false; try { while (true) { + ReadyCallback callback; synchronized (monitor) { if (cursorReturnedDoneOrException) { break; @@ -285,48 +296,61 @@ public void run() { // also stop, even though the callback has not seen the CANCELLED state. cursorReturnedDoneOrException = true; } + callback = AsyncResultSetImpl.this.callback; + } + if (callback == null) { + return; } CallbackResponse response; try { response = callback.cursorReady(AsyncResultSetImpl.this); - } catch (Throwable e) { + } catch (Throwable throwable) { synchronized (monitor) { + resumeRequestedWhileConsuming = false; if (cursorReturnedDoneOrException && state == State.CANCELLED - && e instanceof SpannerException - && ((SpannerException) e).getErrorCode() == ErrorCode.CANCELLED) { + && throwable instanceof SpannerException + && ((SpannerException) throwable).getErrorCode() == ErrorCode.CANCELLED) { // The callback did not catch the cancelled exception (which it should have), but // we'll keep the cancelled state. return; } - executionException = SpannerExceptionFactory.asSpannerException(e); + executionException = SpannerExceptionFactory.asSpannerException(throwable); cursorReturnedDoneOrException = true; } + closeDelegateResultSet(); return; } synchronized (monitor) { if (state == State.CANCELLED) { + resumeRequestedWhileConsuming = false; if (cursorReturnedDoneOrException) { return; } } else { switch (response) { case DONE: + resumeRequestedWhileConsuming = false; state = State.DONE; cursorReturnedDoneOrException = true; return; case PAUSE: + if (resumeRequestedWhileConsuming) { + resumeRequestedWhileConsuming = false; + state = State.RUNNING; + shouldScheduleProducer = true; + return; + } state = State.PAUSED; - // Make sure no-one else is waiting on the current pause latch and create a new - // one. - pausedLatch.countDown(); - pausedLatch = new CountDownLatch(1); return; case CONTINUE: + resumeRequestedWhileConsuming = false; if (buffer.isEmpty()) { - // Call the callback once more if the entire result set has been processed but - // the callback has not yet received a CursorState.DONE or a CANCELLED error. - if (finished && !cursorReturnedDoneOrException) { + // Call the callback once more if the entire result set has been processed or an + // exception was encountered, but the callback has not yet received a + // CursorState.DONE or error. + if ((finished || executionException != null) + && !cursorReturnedDoneOrException) { break; } state = State.RUNNING; @@ -341,17 +365,19 @@ public void run() { } } finally { synchronized (monitor) { - // Count down all latches that the producer might be waiting on. - consumingLatch.countDown(); - while (bufferConsumptionLatch.getCount() > 0L) { - bufferConsumptionLatch.countDown(); - } + callbackRunning = false; + } + if (shouldScheduleProducer) { + scheduleProducerIfNecessary(); } + scheduleCallbackIfNecessary(); + checkCompletion(); } } } private final CallbackRunnable callbackRunnable = new CallbackRunnable(); + private final ProduceRowsRunnable produceRowsRunnable = new ProduceRowsRunnable(); /** * {@link ProduceRowsRunnable} reads data from the underlying {@link ResultSet}, places these in @@ -360,123 +386,243 @@ public void run() { private class ProduceRowsRunnable implements Runnable { @Override public void run() { - boolean stop = false; - boolean hasNext = false; try { - hasNext = delegateResultSet.get().next(); - } catch (Throwable e) { + boolean stopped; synchronized (monitor) { - executionException = SpannerExceptionFactory.asSpannerException(e); - } - } - try { - while (!stop && hasNext) { - try { - synchronized (monitor) { - stop = state.shouldStop; - } - if (!stop) { - while (buffer.remainingCapacity() == 0 && !stop) { - waitIfPaused(); - // The buffer is full and we should let the callback consume a number of rows before - // we proceed with producing any more rows to prevent us from potentially waiting on - // a full buffer repeatedly. - // Wait until at least half of the buffer is available, or if it's a bigger buffer, - // wait until at least 10 rows can be placed in it. - // TODO: Make this more dynamic / configurable? - startCallbackWithBufferLatchIfNecessary( - Math.min( - Math.min(buffer.size() / 2 + 1, buffer.size()), - MAX_WAIT_FOR_BUFFER_CONSUMPTION)); - bufferConsumptionLatch.await(); - synchronized (monitor) { - stop = state.shouldStop; - } - } - } - if (!stop) { - buffer.put(delegateResultSet.get().getCurrentRowAsStruct()); - startCallbackIfNecessary(); - hasNext = delegateResultSet.get().next(); - } - } catch (Throwable e) { - synchronized (monitor) { - executionException = SpannerExceptionFactory.asSpannerException(e); - stop = true; + producerScheduled = false; + stopped = shouldStopProducer(); + if (!stopped) { + if (state == State.STREAMING_INITIALIZED) { + state = State.RUNNING; } + produceRowsInitiated = true; + producerRunning = true; } } - // We don't need any more data from the underlying result set, so we close it as soon as - // possible. Any error that might occur during this will be ignored. - closeDelegateResultSet(); - - // Ensure that the callback has been called at least once, even if the result set was - // cancelled. - synchronized (monitor) { - finished = true; - stop = cursorReturnedDoneOrException; + if (stopped) { + checkCompletion(); + return; } - // Call the callback if there are still rows in the buffer that need to be processed. - while (!stop) { - try { - waitIfPaused(); - startCallbackIfNecessary(); - // Make sure we wait until the callback runner has actually finished. - consumingLatch.await(); - synchronized (monitor) { - stop = cursorReturnedDoneOrException; + while (true) { + synchronized (monitor) { + if (shouldStopProducer() || buffer.remainingCapacity() == 0) { + return; } - } catch (Throwable e) { - result.setException(e); + } + if (!isDataAvailable()) { return; } + + boolean hasNext = delegateResultSet.get().next(); + if (hasNext) { + boolean added = buffer.offer(delegateResultSet.get().getCurrentRowAsStruct()); + Preconditions.checkState(added, "Failed to buffer row despite available capacity"); + scheduleCallbackIfNecessary(); + } else { + synchronized (monitor) { + finished = true; + } + closeDelegateResultSet(); + break; + } } + scheduleCallbackIfNecessary(); + } catch (Throwable throwable) { + setExecutionException(throwable); + scheduleCallbackIfNecessary(); } finally { - if (executorProvider.shouldAutoClose()) { - service.shutdown(); + synchronized (monitor) { + producerRunning = false; } - for (Runnable listener : listeners) { - listener.run(); + scheduleProducerIfNecessary(); + checkCompletion(); + } + } + } + + private boolean isDataAvailable() { + try { + return StreamingUtil.isDataAvailable(delegateResultSet.get()); + } catch (Throwable t) { + return true; + } + } + + private void setExecutionException(Throwable throwable) { + synchronized (monitor) { + if (executionException == null && !state.shouldStop) { + executionException = SpannerExceptionFactory.asSpannerException(throwable); + } + } + } + + private boolean shouldStopProducer() { + return finished + || state.shouldStop + || state == State.PAUSED + || executionException != null + || cursorReturnedDoneOrException; + } + + private boolean canScheduleProducer() { + return !producerRunning + && !producerScheduled + && !shouldStopProducer() + && buffer.remainingCapacity() > 0; + } + + private void scheduleProducerIfNecessary() { + synchronized (monitor) { + if (!canScheduleProducer()) { + return; + } + } + if (!isDataAvailable()) { + return; + } + boolean shouldSchedule = false; + synchronized (monitor) { + if (canScheduleProducer()) { + if (state == State.STREAMING_INITIALIZED) { + state = State.RUNNING; } + produceRowsInitiated = true; + producerScheduled = true; + shouldSchedule = true; + } + } + if (shouldSchedule) { + try { + service.execute(produceRowsRunnable); + } catch (Throwable throwable) { synchronized (monitor) { - if (executionException != null) { - result.setException(executionException); - } else if (state == State.CANCELLED) { - result.setException(CANCELLED_EXCEPTION); - } else { - result.set(null); + producerScheduled = false; + if (executionException == null && !state.shouldStop) { + executionException = SpannerExceptionFactory.asSpannerException(throwable); } } + scheduleCallbackIfNecessary(); + checkCompletion(); } } + } - private void waitIfPaused() throws InterruptedException { - CountDownLatch pause; - synchronized (monitor) { - pause = pausedLatch; + private boolean canScheduleCallback() { + return (state == State.RUNNING || state == State.CANCELLED) + && !cursorReturnedDoneOrException + && !callbackRunning + && !callbackScheduled + && (!buffer.isEmpty() + || finished + || executionException != null + || state == State.CANCELLED); + } + + private void scheduleCallbackIfNecessary() { + boolean shouldExecute = false; + Executor executor = null; + synchronized (monitor) { + if (canScheduleCallback() && this.executor != null) { + if (state == State.RUNNING) { + state = State.CONSUMING; + resumeRequestedWhileConsuming = false; + } + callbackScheduled = true; + shouldExecute = true; + executor = this.executor; } - pause.await(); } - - private void startCallbackIfNecessary() { - startCallbackWithBufferLatchIfNecessary(0); + if (shouldExecute && executor != null) { + try { + executor.execute(callbackRunnable); + } catch (Throwable throwable) { + synchronized (monitor) { + callbackScheduled = false; + cursorReturnedDoneOrException = true; + if (executionException == null && !state.shouldStop) { + executionException = SpannerExceptionFactory.asSpannerException(throwable); + } + } + checkCompletion(); + } } + } - private void startCallbackWithBufferLatchIfNecessary(int bufferLatch) { - synchronized (monitor) { - if ((state == State.RUNNING || state == State.CANCELLED) - && !cursorReturnedDoneOrException) { - consumingLatch = new CountDownLatch(1); - if (bufferLatch > 0) { - bufferConsumptionLatch = new CountDownLatch(bufferLatch); - } - if (state == State.RUNNING) { - state = State.CONSUMING; + private void checkCompletion() { + boolean shouldComplete = false; + synchronized (monitor) { + if (!completed + && !producerRunning + && !producerScheduled + && !callbackRunning + && !callbackScheduled) { + if (state == State.DONE) { + completed = true; + shouldComplete = true; + } else if (state == State.CANCELLED) { + if (cursorReturnedDoneOrException) { + completed = true; + shouldComplete = true; } - executor.execute(callbackRunnable); + } else if (cursorReturnedDoneOrException + && ((finished && buffer.isEmpty()) || executionException != null)) { + state = State.DONE; + completed = true; + shouldComplete = true; } } } + if (shouldComplete) { + cleanupAndCompleteResult(); + } + } + + private void cleanupAndCompleteResult() { + closeDelegateResultSet(); + buffer.clear(); + currentRow = null; + if (executorProvider.shouldAutoClose()) { + service.shutdown(); + } + callback = null; + executor = null; + for (Runnable listener : listeners) { + try { + listener.run(); + } catch (Throwable t) { + log.log(Level.WARNING, "Listener threw exception", t); + } + } + listeners.clear(); + SettableApiFuture resultFuture = this.result; + if (resultFuture != null) { + if (executionException != null) { + resultFuture.setException(executionException); + } else if (state == State.CANCELLED) { + resultFuture.setException(CANCELLED_EXCEPTION); + } else { + resultFuture.set(null); + } + } + } + + @VisibleForTesting + int getBufferSize() { + return buffer.size(); + } + + @VisibleForTesting + State getState() { + synchronized (monitor) { + return state; + } + } + + @VisibleForTesting + boolean isClosed() { + synchronized (monitor) { + return closed; + } } private class InitiateStreamingRunnable implements Runnable { @@ -491,69 +637,118 @@ public void run() { // need to eagerly start the ProduceRowsRunnable. if (!initiateStreaming(AsyncResultSetImpl.this)) { initiateProduceRows(); + } else { + scheduleProducerIfNecessary(); } } catch (Throwable exception) { - executionException = SpannerExceptionFactory.asSpannerException(exception); - initiateProduceRows(); + synchronized (monitor) { + executionException = SpannerExceptionFactory.asSpannerException(exception); + produceRowsInitiated = true; + if (state == State.STREAMING_INITIALIZED) { + state = State.RUNNING; + } + } + scheduleCallbackIfNecessary(); + checkCompletion(); } } } /** Sets the callback for this {@link AsyncResultSet}. */ @Override - public ApiFuture setCallback(Executor exec, ReadyCallback cb) { + public ApiFuture setCallback(Executor executor, ReadyCallback callback) { + SettableApiFuture resultFuture; synchronized (monitor) { Preconditions.checkState(!closed, "This AsyncResultSet has been closed"); Preconditions.checkState( this.state == State.INITIALIZED, "callback may not be set multiple times"); // Start to fetch data and buffer these. - this.result = SettableApiFuture.create(); + this.result = resultFuture = SettableApiFuture.create(); this.state = State.STREAMING_INITIALIZED; + this.executor = MoreExecutors.newSequentialExecutor(Preconditions.checkNotNull(executor)); + this.callback = Preconditions.checkNotNull(callback); + } + try { this.service.execute(new InitiateStreamingRunnable()); - this.executor = MoreExecutors.newSequentialExecutor(Preconditions.checkNotNull(exec)); - this.callback = Preconditions.checkNotNull(cb); - pausedLatch.countDown(); - return result; + } catch (Throwable throwable) { + synchronized (monitor) { + cursorReturnedDoneOrException = true; + if (executionException == null) { + executionException = SpannerExceptionFactory.asSpannerException(throwable); + } + } + checkCompletion(); } + return resultFuture; } private void initiateProduceRows() { synchronized (monitor) { + if (this.produceRowsInitiated) { + return; + } + this.produceRowsInitiated = true; if (this.state == State.STREAMING_INITIALIZED) { this.state = State.RUNNING; } - produceRowsInitiated = true; } - this.service.execute(new ProduceRowsRunnable()); + scheduleProducerIfNecessary(); + checkCompletion(); } - Future getResult() { + @Nullable Future getResult() { return result; } @Override public void cancel() { + boolean shouldStartCallback = false; synchronized (monitor) { Preconditions.checkState( state != State.INITIALIZED && state != State.SYNC, "cannot cancel a result set without a callback"); + if (completed) { + return; + } state = State.CANCELLED; - pausedLatch.countDown(); + resumeRequestedWhileConsuming = false; + if (!callbackRunning && !callbackScheduled) { + shouldStartCallback = true; + } } + closeDelegateResultSet(); + if (shouldStartCallback) { + scheduleCallbackIfNecessary(); + } + checkCompletion(); } @Override public void resume() { + boolean shouldStartCallback = false; + boolean shouldScheduleProducer = false; synchronized (monitor) { Preconditions.checkState( state != State.INITIALIZED && state != State.SYNC, "cannot resume a result set without a callback"); + if (completed) { + return; + } if (state == State.PAUSED) { state = State.RUNNING; - pausedLatch.countDown(); + shouldStartCallback = true; + shouldScheduleProducer = true; + } else if (state == State.CONSUMING) { + resumeRequestedWhileConsuming = true; } } + if (shouldStartCallback) { + scheduleCallbackIfNecessary(); + } + if (shouldScheduleProducer) { + scheduleProducerIfNecessary(); + } } private static class CreateListCallback implements ReadyCallback { @@ -596,10 +791,11 @@ public ApiFuture> toListAsync( Preconditions.checkState(!closed, "This AsyncResultSet has been closed"); Preconditions.checkState( this.state == State.INITIALIZED, "This AsyncResultSet has already been used."); - final SettableApiFuture> res = SettableApiFuture.create(); - CreateListCallback callback = new CreateListCallback<>(res, transformer); + final SettableApiFuture> resultFuture = SettableApiFuture.create(); + CreateListCallback callback = new CreateListCallback<>(resultFuture, transformer); ApiFuture finished = setCallback(executor, callback); - return ApiFutures.transformAsync(finished, ignored -> res, MoreExecutors.directExecutor()); + return ApiFutures.transformAsync( + finished, ignored -> resultFuture, MoreExecutors.directExecutor()); } } @@ -608,10 +804,10 @@ public List toList(Function transformer) throws SpannerE ApiFuture> future = toListAsync(transformer, MoreExecutors.directExecutor()); try { return future.get(); - } catch (ExecutionException e) { - throw SpannerExceptionFactory.asSpannerException(e.getCause()); - } catch (Throwable e) { - throw SpannerExceptionFactory.asSpannerException(e); + } catch (ExecutionException executionException) { + throw SpannerExceptionFactory.asSpannerException(executionException.getCause()); + } catch (Throwable throwable) { + throw SpannerExceptionFactory.asSpannerException(throwable); } } @@ -623,9 +819,9 @@ public boolean next() throws SpannerException { "Cannot call next() on a result set with a callback."); this.state = State.SYNC; } - boolean res = delegateResultSet.get().next(); - currentRow = res ? delegateResultSet.get().getCurrentRowAsStruct() : null; - return res; + boolean hasNext = delegateResultSet.get().next(); + currentRow = hasNext ? delegateResultSet.get().getCurrentRowAsStruct() : null; + return hasNext; } @Override @@ -660,19 +856,16 @@ public Struct getCurrentRowAsStruct() { @Override public void onStreamMessage(PartialResultSet partialResultSet, boolean bufferIsFull) { + boolean shouldInitiate = false; synchronized (monitor) { - if (produceRowsInitiated) { - return; - } - // if PartialResultSet contains a resume token or buffer size is full, or - // we have reached the end of the stream, we can start the thread. - boolean startJobThread = - !partialResultSet.getResumeToken().isEmpty() - || bufferIsFull - || partialResultSet == GrpcStreamIterator.END_OF_STREAM; - if (startJobThread || state != State.STREAMING_INITIALIZED) { - initiateProduceRows(); + if (!produceRowsInitiated) { + shouldInitiate = true; } } + if (shouldInitiate) { + initiateProduceRows(); + } else { + scheduleProducerIfNecessary(); + } } } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ForwardingResultSet.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ForwardingResultSet.java index 3c4883e65862..33f721c90a04 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ForwardingResultSet.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ForwardingResultSet.java @@ -110,4 +110,9 @@ public ResultSetMetadata getMetadata() { public boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMessageListener) { return StreamingUtil.initiateStreaming(delegate.get(), streamMessageListener); } + + @Override + public boolean isDataAvailable() { + return StreamingUtil.isDataAvailable(delegate.get()); + } } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcResultSet.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcResultSet.java index 80a9dfcf5331..7f49fd7c1c3a 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcResultSet.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcResultSet.java @@ -132,6 +132,11 @@ public boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMess return iterator.initiateStreaming(streamMessageListener); } + @Override + public boolean isDataAvailable() { + return iterator.isDataAvailable(); + } + @Override public void close() { synchronized (this) { diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcStreamIterator.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcStreamIterator.java index e0df4c422e3c..10a7d16f05c0 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcStreamIterator.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcStreamIterator.java @@ -47,13 +47,13 @@ class GrpcStreamIterator extends AbstractIterator private final BlockingQueue stream; private final Statement statement; - private SpannerRpc.StreamingCall call; + private volatile SpannerRpc.StreamingCall call; private volatile boolean withBeginTransaction; private final boolean lastStatement; private TimeUnit streamWaitTimeoutUnit; private long streamWaitTimeoutValue; - private SpannerException error; - private boolean done; + private volatile SpannerException error; + private volatile boolean done; @VisibleForTesting GrpcStreamIterator( @@ -129,6 +129,11 @@ public boolean isLastStatement() { return lastStatement; } + @Override + public boolean isDataAvailable() { + return !stream.isEmpty() || error != null || done; + } + @Override protected final PartialResultSet computeNext() { PartialResultSet next; @@ -155,9 +160,11 @@ protected final PartialResultSet computeNext() { call = null; if (error != null) { + done = true; throw SpannerExceptionFactory.asSpannerException(error); } + done = true; endOfData(); return null; } @@ -187,6 +194,7 @@ public void onPartialResultSet(PartialResultSet results) { @Override public void onCompleted() { if (!done) { + done = true; addToStream(END_OF_STREAM); } } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcValueIterator.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcValueIterator.java index 09b850c93f3c..187975802cb3 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcValueIterator.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/GrpcValueIterator.java @@ -42,8 +42,8 @@ private enum StreamValue { private final CloseableIterator stream; private ResultSetMetadata metadata; private Type type; - private PartialResultSet current; - private int pos; + private volatile PartialResultSet current; + private volatile int pos; private ResultSetStats statistics; private final Listener listener; @@ -131,6 +131,11 @@ boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMessageList return stream.initiateStreaming(streamMessageListener); } + boolean isDataAvailable() { + PartialResultSet currentCopy = this.current; + return (currentCopy != null && pos < currentCopy.getValuesCount()) || stream.isDataAvailable(); + } + Type type() { checkState(type != null, "metadata has not been received"); return type; @@ -141,8 +146,9 @@ private boolean ensureReady(StreamValue requiredValue) throws SpannerException { if (!stream.hasNext()) { return false; } - current = stream.next(); + PartialResultSet nextCurrent = stream.next(); pos = 0; + current = nextCurrent; if (type == null) { // This is the first message on the stream. if (!current.hasMetadata() || !current.getMetadata().hasRowType()) { diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ResumableStreamIterator.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ResumableStreamIterator.java index aac7f63c8614..8e6eb86429bf 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ResumableStreamIterator.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/ResumableStreamIterator.java @@ -43,9 +43,10 @@ import java.util.concurrent.CountDownLatch; import java.util.concurrent.Executor; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.logging.Level; import java.util.logging.Logger; -import javax.annotation.Nullable; +import org.jspecify.annotations.Nullable; /** * Wraps an iterator over partial result sets, supporting resuming RPCs on error. This class keeps @@ -64,14 +65,24 @@ abstract class ResumableStreamIterator extends AbstractIterator retryableCodes; private static final Logger logger = Logger.getLogger(ResumableStreamIterator.class.getName()); private BackOff backOff; - private final LinkedList buffer = new LinkedList<>(); + @VisibleForTesting final LinkedList buffer = new LinkedList<>(); private final int maxBufferSize; private final ISpan span; private final TraceWrapper tracer; - private CloseableIterator stream; + + /** Guards state transitions between {@link #stream} and {@link #closed}. */ + private final Object streamLock = new Object(); + + /** Indicates whether this iterator has been closed. */ + private volatile boolean closed; + + /** Ensures the tracing span is finalized at most once. */ + private final AtomicBoolean spanEnded = new AtomicBoolean(false); + + private volatile CloseableIterator stream; private int attempts; private ByteString resumeToken; - private boolean finished; + private volatile boolean finished; private final XGoogSpannerRequestId requestId; /** @@ -79,7 +90,7 @@ abstract class ResumableStreamIterator extends AbstractIterator streamToAssign, boolean requestPrefetch) { + boolean isClosed; + synchronized (streamLock) { + isClosed = this.closed; + if (!isClosed) { + this.stream = streamToAssign; + } + } + if (isClosed) { + if (streamToAssign != null) { + streamToAssign.close(null); + } + endSpanOnce(); + } else if (requestPrefetch && streamToAssign != null) { + streamToAssign.requestPrefetchChunks(); + } + } + + /** Resets the active stream reference during retries or channel failovers. */ + private void resetStream() { + synchronized (streamLock) { + this.stream = null; + } + } + + /** + * Registers the active stream early during {@code startStream()} setup so that in-flight RPCs can + * be cancelled immediately if {@link #close(String)} is invoked before {@code startStream()} + * returns. + */ + void setStream(CloseableIterator stream) { + assignStream(stream, /* requestPrefetch= */ false); + } + @Override public void close(@Nullable String message) { - if (stream != null) { - stream.close(message); - span.end(); - stream = null; + CloseableIterator streamToClose; + synchronized (streamLock) { + closed = true; + streamToClose = this.stream; + this.stream = null; + } + // Perform stream teardown outside the lock to prevent deadlocks with external callbacks. + if (streamToClose != null) { + streamToClose.close(message); + } + endSpanOnce(); + synchronized (buffer) { + buffer.clear(); } } + boolean isClosed() { + return closed; + } + @Override public boolean isWithBeginTransaction() { return stream != null && stream.isWithBeginTransaction(); @@ -246,21 +322,56 @@ public boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMess return true; } + @Override + public boolean isDataAvailable() { + // If closed, return true to unblock any waiting async consumer so it can observe EOF/closure + // instead of hanging indefinitely. + if (closed) { + return true; + } + synchronized (buffer) { + if (!buffer.isEmpty() + && (finished || !safeToRetry || !buffer.getLast().getResumeToken().isEmpty())) { + return true; + } + } + if (finished) { + return true; + } + CloseableIterator currentStream = this.stream; + return currentStream != null && currentStream.isDataAvailable(); + } + @Override protected PartialResultSet computeNext() { int numAttemptsOnOtherChannel = 0; Context context = Context.current(); - while (true) { + while (!closed) { // Eagerly start stream before consuming any buffered items. startGrpcStreaming(); + if (closed) { + break; + } // Buffer contains items up to a resume token or has reached capacity: flush. - if (!buffer.isEmpty() - && (finished || !safeToRetry || !buffer.getLast().getResumeToken().isEmpty())) { - return buffer.pop(); + PartialResultSet buffered = null; + synchronized (buffer) { + if (!buffer.isEmpty() + && (finished || !safeToRetry || !buffer.getLast().getResumeToken().isEmpty())) { + buffered = buffer.pop(); + } + } + if (buffered != null) { + return buffered; + } + // Snapshot the volatile stream reference to guard against concurrent close() nulling the + // field between check and dereference. + CloseableIterator currentStream = this.stream; + if (currentStream == null) { + break; } try { - if (stream.hasNext()) { - PartialResultSet next = stream.next(); + if (currentStream.hasNext()) { + PartialResultSet next = currentStream.next(); boolean hasResumeToken = !next.getResumeToken().isEmpty(); if (hasResumeToken) { resumeToken = next.getResumeToken(); @@ -269,20 +380,23 @@ protected PartialResultSet computeNext() { // If the buffer is empty and this chunk has a resume token or we cannot resume safely // anyway, we can yield it immediately rather than placing it in the buffer to be // returned on the next iteration. - if ((hasResumeToken || !safeToRetry) && buffer.isEmpty()) { - return next; - } - buffer.add(next); - if (buffer.size() > maxBufferSize && buffer.getLast().getResumeToken().isEmpty()) { - // We need to flush without a restart token. Errors encountered until we see - // such a token will fail the read. - safeToRetry = false; + synchronized (buffer) { + if ((hasResumeToken || !safeToRetry) && buffer.isEmpty()) { + return next; + } + buffer.add(next); + if (buffer.size() > maxBufferSize && buffer.getLast().getResumeToken().isEmpty()) { + // We need to flush without a restart token. Errors encountered until we see + // such a token will fail the read. + safeToRetry = false; + } } } else { finished = true; - if (buffer.isEmpty()) { - endOfData(); - return null; + synchronized (buffer) { + if (buffer.isEmpty()) { + break; + } } } } catch (SpannerException spannerException) { @@ -290,12 +404,14 @@ protected PartialResultSet computeNext() { span.addAnnotation("Stream broken. Safe to retry", spannerException); logger.log(Level.FINE, "Retryable exception, will sleep and retry", spannerException); // Truncate any items in the buffer before the last retry token. - while (!buffer.isEmpty() && buffer.getLast().getResumeToken().isEmpty()) { - buffer.removeLast(); + synchronized (buffer) { + while (!buffer.isEmpty() && buffer.getLast().getResumeToken().isEmpty()) { + buffer.removeLast(); + } + assert buffer.isEmpty() || buffer.getLast().getResumeToken().equals(resumeToken); } - assert buffer.isEmpty() || buffer.getLast().getResumeToken().equals(resumeToken); - stream = null; - try (IScope s = tracer.withSpan(span)) { + resetStream(); + try (IScope scope = tracer.withSpan(span)) { long delay = spannerException.getRetryDelayInMillis(); if (delay != -1) { backoffSleep(context, delay); @@ -310,12 +426,16 @@ protected PartialResultSet computeNext() { continue; } // Check if we should retry the request on a different gRPC channel. - if (resumeToken == null && buffer.isEmpty()) { + boolean bufferIsEmpty; + synchronized (buffer) { + bufferIsEmpty = buffer.isEmpty(); + } + if (resumeToken == null && bufferIsEmpty) { Throwable translated = errorHandler.translateException(spannerException); if (translated instanceof RetryOnDifferentGrpcChannelException) { if (++numAttemptsOnOtherChannel < errorHandler.getMaxAttempts() && prepareIteratorForRetryOnDifferentGrpcChannel()) { - stream = null; + resetStream(); continue; } } @@ -329,21 +449,31 @@ && prepareIteratorForRetryOnDifferentGrpcChannel()) { throw e; } } + endOfData(); + return null; } + /** + * Lazily starts the underlying gRPC stream under {@link #streamLock} if streaming has not already + * been initiated and the iterator has not been closed. + */ private void startGrpcStreaming() { - if (stream == null) { - span.addAnnotation( - "Starting/Resuming stream", - "ResumeToken", - resumeToken == null ? "null" : resumeToken.toStringUtf8()); - try (IScope scope = tracer.withSpan(span)) { - // When start a new stream set the Span as current to make the gRPC Span a child of - // this Span. - stream = checkNotNull(startStream(resumeToken, streamMessageListener, requestId)); - stream.requestPrefetchChunks(); + synchronized (streamLock) { + if (stream != null || closed) { + return; } } + span.addAnnotation( + "Starting/Resuming stream", + "ResumeToken", + resumeToken == null ? "null" : resumeToken.toStringUtf8()); + try (IScope scope = tracer.withSpan(span)) { + // When start a new stream set the Span as current to make the gRPC Span a child of + // this Span. + CloseableIterator streamIterator = + checkNotNull(startStream(resumeToken, streamMessageListener, requestId)); + assignStream(streamIterator, /* requestPrefetch= */ true); + } } boolean isRetryable(SpannerException spannerException) { diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingResultSet.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingResultSet.java index 47b10d852c64..1d23d4a7e2ec 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingResultSet.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingResultSet.java @@ -28,4 +28,13 @@ interface StreamingResultSet extends ResultSet { */ @InternalApi boolean initiateStreaming(AsyncResultSet.StreamMessageListener streamMessageListener); + + /** + * Returns true if data (a chunk, row, EOF, or error) is available to read immediately without + * blocking the calling thread on network I/O. + */ + @InternalApi + default boolean isDataAvailable() { + return true; + } } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingUtil.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingUtil.java index 54496d39f965..54948f949f3f 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingUtil.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/StreamingUtil.java @@ -27,4 +27,15 @@ static boolean initiateStreaming( } return false; } + + static boolean isDataAvailable(ResultSet resultSet) { + ResultSet delegate = resultSet; + while (delegate instanceof ForwardingResultSet) { + delegate = ((ForwardingResultSet) delegate).getDelegate(); + } + if (delegate instanceof StreamingResultSet) { + return ((StreamingResultSet) delegate).isDataAvailable(); + } + return true; + } } diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/AsyncResultSetImplTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/AsyncResultSetImplTest.java index 23180356cbfa..76532c852d00 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/AsyncResultSetImplTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/AsyncResultSetImplTest.java @@ -20,15 +20,20 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertNull; import static org.junit.Assert.assertThrows; +import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; import static org.mockito.Mockito.withSettings; import com.google.api.core.ApiFuture; +import com.google.api.core.SettableApiFuture; import com.google.api.gax.core.ExecutorProvider; import com.google.cloud.spanner.AsyncResultSet.CallbackResponse; import com.google.cloud.spanner.AsyncResultSet.CursorState; @@ -38,18 +43,28 @@ import com.google.protobuf.ByteString; import com.google.protobuf.Value; import com.google.spanner.v1.PartialResultSet; +import com.google.spanner.v1.ResultSetMetadata; +import com.google.spanner.v1.ResultSetStats; +import java.util.ArrayList; +import java.util.Collections; import java.util.List; +import java.util.Random; import java.util.concurrent.BlockingDeque; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutionException; import java.util.concurrent.Executor; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.Future; import java.util.concurrent.LinkedBlockingDeque; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.Semaphore; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.Assume; import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; @@ -520,6 +535,7 @@ public void testOnStreamMessageWhenResumeTokenIsPresent() { Mockito.when( delegate.initiateStreaming(Mockito.any(AsyncResultSet.StreamMessageListener.class))) .thenReturn(true); + Mockito.when(delegate.isDataAvailable()).thenReturn(true); rs.setCallback(Executors.newSingleThreadExecutor(), ignored -> CallbackResponse.DONE); rs.onStreamMessage( @@ -541,6 +557,7 @@ public void testOnStreamMessageWhenCurrentBufferSizeReachedPrefetchChunkSize() { Mockito.when( delegate.initiateStreaming(Mockito.any(AsyncResultSet.StreamMessageListener.class))) .thenReturn(true); + Mockito.when(delegate.isDataAvailable()).thenReturn(true); rs.setCallback(Executors.newSingleThreadExecutor(), ignored -> CallbackResponse.DONE); rs.onStreamMessage( @@ -563,7 +580,1635 @@ public void testOnStreamMessageWhenAsyncResultIsCancelled() { rs.cancel(); rs.onStreamMessage( PartialResultSet.newBuilder().addValues(Value.newBuilder().build()).build(), false); - Mockito.verify(mockedProvider.getExecutor(), times(2)).execute(Mockito.any()); + Mockito.verify(mockedProvider.getExecutor(), times(1)).execute(Mockito.any()); + } + } + + @Test + public void testSequentialRowDeliveryUnderConcurrentEvents() throws Exception { + ExecutorService callbackExecutor = Executors.newFixedThreadPool(2); + ExecutorService resumeExecutor = Executors.newFixedThreadPool(2); + final int rowCount = 100; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Boolean answer(InvocationOnMock invocation) { + currentRowIndex++; + return currentRowIndex <= rowCount; + } + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Struct answer(InvocationOnMock invocation) { + currentRowIndex++; + return Struct.newBuilder().set("ID").to((long) currentRowIndex).build(); + } + }); + + final List receivedRowIds = Collections.synchronizedList(new ArrayList<>()); + final CountDownLatch finishedLatch = new CountDownLatch(1); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState cursorState; + while ((cursorState = resultSet.tryNext()) == CursorState.OK) { + receivedRowIds.add(resultSet.getLong("ID")); + } + if (cursorState == CursorState.DONE) { + finishedLatch.countDown(); + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + final AtomicBoolean testRunning = new AtomicBoolean(true); + resumeExecutor.execute( + () -> { + while (testRunning.get()) { + asyncResultSet.resume(); + Thread.yield(); + } + }); + + assertTrue(finishedLatch.await(10, TimeUnit.SECONDS)); + testRunning.set(false); + assertNull(callbackFuture.get(5, TimeUnit.SECONDS)); + + assertEquals(rowCount, receivedRowIds.size()); + for (int i = 0; i < rowCount; i++) { + assertEquals((long) (i + 1), (long) receivedRowIds.get(i)); + } + } finally { + callbackExecutor.shutdown(); + resumeExecutor.shutdown(); + } + } + + @Test + public void testNonStreamingResultSetWithSmallBufferCapacity() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + final int rowCount = 50; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Boolean answer(InvocationOnMock invocation) { + currentRowIndex++; + return currentRowIndex <= rowCount; + } + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Struct answer(InvocationOnMock invocation) { + currentRowIndex++; + return Struct.newBuilder().set("ID").to((long) currentRowIndex).build(); + } + }); + + final List receivedRowIds = new ArrayList<>(); + // Use buffer size = 1 to force maximum contention and buffer exhaustion + try (AsyncResultSetImpl asyncResultSet = new AsyncResultSetImpl(simpleProvider, delegate, 1)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState cursorState; + while ((cursorState = resultSet.tryNext()) == CursorState.OK) { + receivedRowIds.add(resultSet.getLong("ID")); + } + if (cursorState == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + assertNull(callbackFuture.get(10, TimeUnit.SECONDS)); + assertEquals(rowCount, receivedRowIds.size()); + for (int i = 0; i < rowCount; i++) { + assertEquals((long) (i + 1), (long) receivedRowIds.get(i)); + } + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testPostStreamDrainingWithSlowConsumer() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + final int rowCount = 5; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Boolean answer(InvocationOnMock invocation) { + currentRowIndex++; + return currentRowIndex <= rowCount; + } + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Struct answer(InvocationOnMock invocation) { + currentRowIndex++; + return Struct.newBuilder().set("ID").to((long) currentRowIndex).build(); + } + }); + + final List receivedRowIds = new ArrayList<>(); + final AtomicBoolean listenerInvoked = new AtomicBoolean(false); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + asyncResultSet.addListener(() -> listenerInvoked.set(true)); + + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState cursorState = resultSet.tryNext(); + if (cursorState == CursorState.OK) { + receivedRowIds.add(resultSet.getLong("ID")); + return CallbackResponse.PAUSE; + } else if (cursorState == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + while (!callbackFuture.isDone()) { + Thread.sleep(10); + asyncResultSet.resume(); + } + + assertNull(callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(listenerInvoked.get()); + assertEquals(rowCount, receivedRowIds.size()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testStreamingResultSetWithMultipleRowsPerChunk() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + StreamingResultSet delegate = mock(StreamingResultSet.class); + when(delegate.isDataAvailable()).thenReturn(true); + final int rowCount = 10; + when(delegate.next()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Boolean answer(InvocationOnMock invocation) { + currentRowIndex++; + return currentRowIndex <= rowCount; + } + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Struct answer(InvocationOnMock invocation) { + currentRowIndex++; + return Struct.newBuilder().set("ID").to((long) currentRowIndex).build(); + } + }); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + when(delegate.initiateStreaming(any(AsyncResultSet.StreamMessageListener.class))) + .thenAnswer( + answer -> { + AsyncResultSet.StreamMessageListener listener = answer.getArgument(0); + // Deliver one chunk containing a resume token representing all rows + listener.onStreamMessage( + PartialResultSet.newBuilder() + .setResumeToken(ByteString.copyFromUtf8("resume-token")) + .build(), + false); + return true; + }); + + final List receivedRowIds = new ArrayList<>(); + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState cursorState; + while ((cursorState = resultSet.tryNext()) == CursorState.OK) { + receivedRowIds.add(resultSet.getLong("ID")); + } + if (cursorState == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + assertNull(callbackFuture.get(5, TimeUnit.SECONDS)); + assertEquals(rowCount, receivedRowIds.size()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testRapidConcurrentPauseResume() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ExecutorService resumeExecutor = Executors.newFixedThreadPool(4); + final int rowCount = 100; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Boolean answer(InvocationOnMock invocation) { + currentRowIndex++; + return currentRowIndex <= rowCount; + } + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + new Answer() { + int currentRowIndex = 0; + + @Override + public Struct answer(InvocationOnMock invocation) { + currentRowIndex++; + return Struct.newBuilder().set("ID").to((long) currentRowIndex).build(); + } + }); + + final List receivedRowIds = Collections.synchronizedList(new ArrayList<>()); + final Random random = new Random(); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState cursorState = resultSet.tryNext(); + if (cursorState == CursorState.OK) { + receivedRowIds.add(resultSet.getLong("ID")); + return random.nextBoolean() ? CallbackResponse.PAUSE : CallbackResponse.CONTINUE; + } else if (cursorState == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + final AtomicBoolean testRunning = new AtomicBoolean(true); + for (int i = 0; i < 4; i++) { + resumeExecutor.execute( + () -> { + while (testRunning.get()) { + asyncResultSet.resume(); + Thread.yield(); + } + }); + } + + assertNull(callbackFuture.get(10, TimeUnit.SECONDS)); + testRunning.set(false); + + assertEquals(rowCount, receivedRowIds.size()); + for (int i = 0; i < rowCount; i++) { + assertEquals((long) (i + 1), (long) receivedRowIds.get(i)); + } + } finally { + callbackExecutor.shutdown(); + resumeExecutor.shutdown(); + } + } + + @Test + public void testErrorAfterBufferedRowsDelivered() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenReturn(true) + .thenReturn(true) + .thenReturn(true) + .thenThrow( + SpannerExceptionFactory.newSpannerException( + ErrorCode.UNAVAILABLE, "temporary network glitch")); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + final List receivedStates = new ArrayList<>(); + final AtomicInteger rowCount = new AtomicInteger(); + final SettableApiFuture exceptionFuture = SettableApiFuture.create(); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + try { + CursorState cursorState = resultSet.tryNext(); + receivedStates.add(cursorState); + if (cursorState == CursorState.OK) { + rowCount.incrementAndGet(); + return CallbackResponse.CONTINUE; + } + return CallbackResponse.DONE; + } catch (SpannerException e) { + exceptionFuture.set(e); + return CallbackResponse.DONE; + } + }); + + SpannerException callbackException = exceptionFuture.get(5, TimeUnit.SECONDS); + assertEquals(ErrorCode.UNAVAILABLE, callbackException.getErrorCode()); + assertTrue(callbackException.getMessage().contains("temporary network glitch")); + assertEquals(3, rowCount.get()); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertEquals( + ErrorCode.UNAVAILABLE, ((SpannerException) executionException.getCause()).getErrorCode()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testProducerYieldsThreadWhenWaitingForStreamData() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + GrpcResultSet delegate = mock(GrpcResultSet.class); + AtomicReference listenerRef = new AtomicReference<>(); + CountDownLatch listenerRegistered = new CountDownLatch(1); + when(delegate.initiateStreaming(any(AsyncResultSet.StreamMessageListener.class))) + .thenAnswer( + answer -> { + listenerRef.set(answer.getArgument(0)); + listenerRegistered.countDown(); + return true; + }); + final AtomicInteger nextCallCount = new AtomicInteger(); + when(delegate.next()) + .thenAnswer( + invocation -> { + int count = nextCallCount.incrementAndGet(); + return count <= 2; + }); + when(delegate.getCurrentRowAsStruct()) + .thenAnswer( + invocation -> Struct.newBuilder().set("ID").to((long) nextCallCount.get()).build()); + + // Initially isDataAvailable is true for the first row, then false until chunk 2 arrives + AtomicBoolean secondChunkAvailable = new AtomicBoolean(false); + when(delegate.isDataAvailable()) + .thenAnswer(invocation -> nextCallCount.get() == 0 || secondChunkAvailable.get()); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + List receivedRows = new ArrayList<>(); + CountDownLatch firstRowReceived = new CountDownLatch(1); + + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state; + while ((state = resultSet.tryNext()) == CursorState.OK) { + receivedRows.add(resultSet.getLong("ID")); + firstRowReceived.countDown(); + } + if (state == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + // Wait for the streaming listener to be registered asynchronously + assertTrue(listenerRegistered.await(5, TimeUnit.SECONDS)); + + // Deliver first message to trigger the producer + listenerRef + .get() + .onStreamMessage( + PartialResultSet.newBuilder() + .setResumeToken(ByteString.copyFromUtf8("chunk-1")) + .build(), + false); + assertTrue(firstRowReceived.await(5, TimeUnit.SECONDS)); + assertEquals(1, receivedRows.size()); + + // While stream is waiting for chunk 2, delegate.isDataAvailable() is false. + // Submit another task to simpleProvider (which is a 1-thread pool!). + // If the producer did not yield, the 1-thread pool would be blocked and this task could NOT + // execute! + ScheduledExecutorService producerExecutor = simpleProvider.getExecutor(); + Future pingTask = producerExecutor.submit(() -> "pong"); + assertEquals("pong", pingTask.get(2, TimeUnit.SECONDS)); + + // Now make chunk 2 available and notify listener + secondChunkAvailable.set(true); + listenerRef + .get() + .onStreamMessage( + PartialResultSet.newBuilder() + .setResumeToken(ByteString.copyFromUtf8("chunk-2")) + .build(), + false); + + assertNull(callbackFuture.get(5, TimeUnit.SECONDS)); + assertEquals(2, receivedRows.size()); + assertEquals(1L, (long) receivedRows.get(0)); + assertEquals(2L, (long) receivedRows.get(1)); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testStreamInitializationFailurePropagatesToCallbackAndFuture() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + StreamingResultSet delegate = mock(StreamingResultSet.class); + when(delegate.initiateStreaming(any(AsyncResultSet.StreamMessageListener.class))) + .thenThrow( + SpannerExceptionFactory.newSpannerException( + ErrorCode.INVALID_ARGUMENT, "Invalid query syntax error")); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + final SettableApiFuture exceptionFromCallback = SettableApiFuture.create(); + + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + try { + resultSet.tryNext(); + } catch (SpannerException exception) { + exceptionFromCallback.set(exception); + } + return CallbackResponse.DONE; + }); + + SpannerException callbackException = exceptionFromCallback.get(5, TimeUnit.SECONDS); + assertEquals(ErrorCode.INVALID_ARGUMENT, callbackException.getErrorCode()); + assertTrue(callbackException.getMessage().contains("Invalid query syntax error")); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertEquals( + ErrorCode.INVALID_ARGUMENT, + ((SpannerException) executionException.getCause()).getErrorCode()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testCancelDoesNotRunListenersUntilCallbackFinishes() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(Struct.newBuilder().set("ID").to(1L).build()); + + final CountDownLatch callbackStartedLatch = new CountDownLatch(1); + final CountDownLatch allowCallbackToFinishLatch = new CountDownLatch(1); + final AtomicBoolean listenerRanWhileCallbackExecuting = new AtomicBoolean(false); + final AtomicBoolean listenerRan = new AtomicBoolean(false); + final AtomicBoolean callbackFinished = new AtomicBoolean(false); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + asyncResultSet.addListener( + () -> { + listenerRan.set(true); + if (!callbackFinished.get()) { + listenerRanWhileCallbackExecuting.set(true); + } + }); + + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state = resultSet.tryNext(); + if (state == CursorState.OK) { + callbackStartedLatch.countDown(); + try { + assertTrue(allowCallbackToFinishLatch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException interruptedException) { + Thread.currentThread().interrupt(); + } + callbackFinished.set(true); + return CallbackResponse.CONTINUE; + } + return CallbackResponse.DONE; + }); + + // Wait until the callback is actively executing inside cursorReady + assertTrue(callbackStartedLatch.await(5, TimeUnit.SECONDS)); + + // Cancel while the callback is still blocked inside cursorReady + asyncResultSet.cancel(); + + // Ensure that listeners did NOT run while the callback was still executing + assertFalse(listenerRan.get()); + assertFalse(listenerRanWhileCallbackExecuting.get()); + + // Release callback so it can complete + allowCallbackToFinishLatch.countDown(); + + // Now verify that the callback future completes with CANCELLED + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertEquals( + ErrorCode.CANCELLED, ((SpannerException) executionException.getCause()).getErrorCode()); + + // And the listener has run AFTER callback completed + assertTrue(listenerRan.get()); + assertFalse(listenerRanWhileCallbackExecuting.get()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testGrpcStreamIteratorDoneSetOnCompleted() { + GrpcStreamIterator streamIterator = new GrpcStreamIterator(false, 4, false); + assertFalse(streamIterator.isDataAvailable()); + + streamIterator.consumer().onCompleted(); + assertTrue(streamIterator.isDataAvailable()); + assertFalse(streamIterator.hasNext()); + assertTrue(streamIterator.isDataAvailable()); + } + + @Test + public void testBufferClearedOnCancellation() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, true, true, false); + when(delegate.getCurrentRowAsStruct()) + .thenReturn( + Struct.newBuilder().set("ID").to(1L).build(), + Struct.newBuilder().set("ID").to(2L).build(), + Struct.newBuilder().set("ID").to(3L).build()); + + final CountDownLatch callbackStartedLatch = new CountDownLatch(1); + final CountDownLatch allowCancellationLatch = new CountDownLatch(1); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state = resultSet.tryNext(); + if (state == CursorState.OK) { + callbackStartedLatch.countDown(); + try { + assertTrue(allowCancellationLatch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException interruptedException) { + Thread.currentThread().interrupt(); + } + return CallbackResponse.CONTINUE; + } + return CallbackResponse.DONE; + }); + + assertTrue(callbackStartedLatch.await(5, TimeUnit.SECONDS)); + asyncResultSet.cancel(); + allowCancellationLatch.countDown(); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertEquals( + ErrorCode.CANCELLED, ((SpannerException) executionException.getCause()).getErrorCode()); + assertEquals(0, asyncResultSet.getBufferSize()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testForwardingResultSetIsDataAvailable() { + GrpcResultSet grpcDelegate = mock(GrpcResultSet.class); + when(grpcDelegate.isDataAvailable()).thenReturn(false, true); + + ForwardingResultSet innerForwarding = new ForwardingResultSet(grpcDelegate); + ForwardingResultSet outerForwarding = new ForwardingResultSet(innerForwarding); + + assertFalse(outerForwarding.isDataAvailable()); + assertTrue(outerForwarding.isDataAvailable()); + } + + @Test + public void testProducerStopsImmediatelyWhenCallbackThrowsUncheckedException() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state = resultSet.tryNext(); + if (state == CursorState.OK) { + throw new RuntimeException("User callback error"); + } + return CallbackResponse.CONTINUE; + }); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> future.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertTrue(executionException.getCause().getMessage().contains("User callback error")); + assertEquals(0, asyncResultSet.getBufferSize()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testRejectedExecutionExceptionOnProducerService() throws Exception { + ScheduledExecutorService failingExecutor = mock(ScheduledExecutorService.class); + Mockito.doThrow(new RejectedExecutionException("Producer rejected")) + .when(failingExecutor) + .execute(any(Runnable.class)); + ExecutorProvider failingProvider = + new ExecutorProvider() { + @Override + public boolean shouldAutoClose() { + return false; + } + + @Override + public ScheduledExecutorService getExecutor() { + return failingExecutor; + } + }; + + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl( + failingProvider, mock(ResultSet.class), AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback(callbackExecutor, resultSet -> CallbackResponse.DONE); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> future.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertTrue(executionException.getCause().getMessage().contains("Producer rejected")); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testRejectedExecutionExceptionOnCallbackExecutor() { + Executor failingCallbackExecutor = + command -> { + throw new RejectedExecutionException("Callback rejected"); + }; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback(failingCallbackExecutor, resultSet -> CallbackResponse.DONE); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> future.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertTrue(executionException.getCause().getMessage().contains("Callback rejected")); + } + } + + @Test + public void testIsUsed() throws Exception { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + assertFalse(asyncResultSet.isUsed()); + ExecutorService executor = Executors.newSingleThreadExecutor(); + try { + ApiFuture future = + asyncResultSet.setCallback( + executor, + resultSet -> { + resultSet.tryNext(); + return CallbackResponse.DONE; + }); + assertTrue(asyncResultSet.isUsed()); + future.get(5, TimeUnit.SECONDS); + assertTrue(asyncResultSet.isUsed()); + } finally { + executor.shutdown(); + } + } + + ResultSet delegateSync = mock(ResultSet.class); + when(delegateSync.next()).thenReturn(true, false); + when(delegateSync.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + try (AsyncResultSetImpl asyncResultSetSync = + new AsyncResultSetImpl( + simpleProvider, delegateSync, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + assertFalse(asyncResultSetSync.isUsed()); + assertTrue(asyncResultSetSync.next()); + assertTrue(asyncResultSetSync.isUsed()); + } + } + + @Test + public void testListenerPreconditionsAndRemoval() throws Exception { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(false); + + AtomicBoolean removedListenerCalled = new AtomicBoolean(false); + AtomicBoolean activeListenerCalled = new AtomicBoolean(false); + Runnable removedListener = () -> removedListenerCalled.set(true); + Runnable activeListener = () -> activeListenerCalled.set(true); + + ExecutorService executor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + asyncResultSet.addListener(removedListener); + asyncResultSet.addListener(activeListener); + asyncResultSet.removeListener(removedListener); + + ApiFuture future = + asyncResultSet.setCallback( + executor, + resultSet -> { + resultSet.tryNext(); + return CallbackResponse.DONE; + }); + + assertThrows(IllegalStateException.class, () -> asyncResultSet.addListener(() -> {})); + assertThrows( + IllegalStateException.class, () -> asyncResultSet.removeListener(activeListener)); + + future.get(5, TimeUnit.SECONDS); + assertFalse(removedListenerCalled.get()); + assertTrue(activeListenerCalled.get()); + } finally { + executor.shutdown(); + } + } + + @Test + public void testListenerExceptionDoesNotPreventSubsequentListenersOrCompletion() + throws Exception { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(false); + + AtomicBoolean secondListenerCalled = new AtomicBoolean(false); + ExecutorService executor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + asyncResultSet.addListener( + () -> { + throw new RuntimeException("Listener failure"); + }); + asyncResultSet.addListener(() -> secondListenerCalled.set(true)); + + ApiFuture future = + asyncResultSet.setCallback( + executor, + resultSet -> { + resultSet.tryNext(); + return CallbackResponse.DONE; + }); + + assertNull(future.get(5, TimeUnit.SECONDS)); + assertTrue(secondListenerCalled.get()); + } finally { + executor.shutdown(); + } + } + + @Test + public void testGetStatsAndGetMetadata() { + ResultSet delegate = mock(ResultSet.class); + ResultSetStats expectedStats = ResultSetStats.getDefaultInstance(); + ResultSetMetadata expectedMetadata = ResultSetMetadata.getDefaultInstance(); + when(delegate.getStats()).thenReturn(expectedStats); + when(delegate.getMetadata()).thenReturn(expectedMetadata); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + assertEquals(expectedStats, asyncResultSet.getStats()); + assertEquals(expectedMetadata, asyncResultSet.getMetadata()); + } + } + + @Test + public void testSyncModeOperations() { + ResultSet delegate = mock(ResultSet.class); + Struct mockRow = mock(Struct.class); + when(delegate.next()).thenReturn(true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(mockRow); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + assertTrue(asyncResultSet.next()); + assertEquals(mockRow, asyncResultSet.getCurrentRowAsStruct()); + + assertThrows( + IllegalStateException.class, + () -> + asyncResultSet.setCallback( + Executors.newSingleThreadExecutor(), resultSet -> CallbackResponse.DONE)); + + assertFalse(asyncResultSet.next()); + assertNull(asyncResultSet.getCurrentRowAsStruct()); + + ResultSet anotherDelegate = mock(ResultSet.class); + try (AsyncResultSetImpl asyncResultSet2 = + new AsyncResultSetImpl( + simpleProvider, anotherDelegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + asyncResultSet2.setCallback( + Executors.newSingleThreadExecutor(), resultSet -> CallbackResponse.DONE); + assertThrows(IllegalStateException.class, asyncResultSet2::next); + } + + asyncResultSet.close(); + Mockito.verify(delegate, times(1)).close(); + assertThrows(IllegalStateException.class, asyncResultSet::getCurrentRowAsStruct); + } + } + + @Test + public void testExecutorProviderShouldAutoClose() throws Exception { + ScheduledExecutorService realService = Executors.newSingleThreadScheduledExecutor(); + ExecutorProvider autoCloseProvider = + new ExecutorProvider() { + @Override + public boolean shouldAutoClose() { + return true; + } + + @Override + public ScheduledExecutorService getExecutor() { + return realService; + } + }; + + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(false); + + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl( + autoCloseProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + resultSet.tryNext(); + return CallbackResponse.DONE; + }); + + future.get(5, TimeUnit.SECONDS); + assertTrue(realService.isShutdown()); + } finally { + callbackExecutor.shutdown(); + realService.shutdown(); + } + } + + @Test + public void testDefaultIsDataAvailable() { + StreamingResultSet defaultStreamingResultSet = + mock(StreamingResultSet.class, Mockito.CALLS_REAL_METHODS); + assertTrue(defaultStreamingResultSet.isDataAvailable()); + + AbstractResultSet.CloseableIterator defaultIterator = + mock(AbstractResultSet.CloseableIterator.class, Mockito.CALLS_REAL_METHODS); + assertTrue(defaultIterator.isDataAvailable()); + } + + @Test + public void testDuplicateCloseIsIdempotent() { + ResultSet delegate = mock(ResultSet.class); + AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE); + asyncResultSet.close(); + verify(delegate, times(1)).close(); + // Second close is a no-op + asyncResultSet.close(); + verify(delegate, times(1)).close(); + } + + @Test + public void testCallbackThrowsCancelledException() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + asyncResultSet.cancel(); + // Calling tryNext() will throw CANCELLED_EXCEPTION out of cursorReady + resultSet.tryNext(); + return CallbackResponse.CONTINUE; + }); + SpannerException exception = assertThrows(SpannerException.class, () -> get(future)); + assertEquals(ErrorCode.CANCELLED, exception.getErrorCode()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testSetCallbackInitiateStreamingRejection() { + ScheduledExecutorService rejectingExecutor = mock(ScheduledExecutorService.class); + Mockito.doThrow(new RejectedExecutionException("rejected")) + .when(rejectingExecutor) + .execute(any(Runnable.class)); + ExecutorProvider rejectingProvider = + new ExecutorProvider() { + @Override + public boolean shouldAutoClose() { + return false; + } + + @Override + public ScheduledExecutorService getExecutor() { + return rejectingExecutor; + } + }; + + ResultSet delegate = mock(ResultSet.class); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl( + rejectingProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + Executors.newSingleThreadExecutor(), resultSet -> CallbackResponse.DONE); + SpannerException exception = assertThrows(SpannerException.class, () -> get(future)); + assertTrue(exception.getMessage().contains("rejected")); + } + } + + @Test + public void testIsDataAvailableExceptionHandling() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + StreamingResultSet delegate = mock(StreamingResultSet.class); + when(delegate.isDataAvailable()).thenThrow(new RuntimeException("isDataAvailable error")); + when(delegate.next()).thenReturn(false); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + assertEquals(CursorState.DONE, resultSet.tryNext()); + return CallbackResponse.DONE; + }); + assertNull(future.get(5, TimeUnit.SECONDS)); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testCallbackProcessingOneRowAtATimeReceivesDone() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, true, false); + when(delegate.getCurrentRowAsStruct()) + .thenReturn( + Struct.newBuilder().set("v").to("a").build(), + Struct.newBuilder().set("v").to("b").build()); + + AtomicInteger rowsSeen = new AtomicInteger(); + AtomicBoolean doneReceived = new AtomicBoolean(); + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state = resultSet.tryNext(); + if (state == CursorState.OK) { + rowsSeen.incrementAndGet(); + return CallbackResponse.CONTINUE; + } + if (state == CursorState.DONE) { + doneReceived.set(true); + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + assertNull(future.get(5, TimeUnit.SECONDS)); + assertTrue(doneReceived.get()); + assertEquals(2, rowsSeen.get()); + } finally { + callbackExecutor.shutdown(); + } + } + + @Test + public void testHighVolumeStreamingWithoutBufferStarvation() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + int totalRowCount = 50_000; + int bufferCapacity = 50; + + StreamingResultSet delegate = + new ForwardingResultSet(mock(ResultSet.class)) { + private int currentRow = 0; + private Struct currentStruct; + + @Override + public boolean next() { + if (currentRow < totalRowCount) { + currentRow++; + currentStruct = Struct.newBuilder().set("ID").to((long) currentRow).build(); + return true; + } + return false; + } + + @Override + public Struct getCurrentRowAsStruct() { + return currentStruct; + } + + @Override + public boolean isDataAvailable() { + return true; + } + + @Override + public boolean initiateStreaming(AsyncResultSet.StreamMessageListener listener) { + return true; + } + }; + + AtomicInteger rowsReceived = new AtomicInteger(); + AtomicBoolean doneReceived = new AtomicBoolean(); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, bufferCapacity)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + while (true) { + switch (resultSet.tryNext()) { + case OK: + int expectedId = rowsReceived.incrementAndGet(); + assertEquals(expectedId, resultSet.getLong("ID")); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + doneReceived.set(true); + return CallbackResponse.DONE; + } + } + }); + + assertNull(future.get(30, TimeUnit.SECONDS)); + assertTrue(doneReceived.get()); + assertEquals(totalRowCount, rowsReceived.get()); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testHighVolumeStreamingWithAsyncChunkArrivals() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + ExecutorService feederExecutor = Executors.newSingleThreadExecutor(); + int totalRowCount = 10_000; + int chunkSize = 100; + int bufferCapacity = 50; + + AtomicInteger availableRows = new AtomicInteger(); + AtomicReference registeredListener = + new AtomicReference<>(); + CountDownLatch streamingInitiated = new CountDownLatch(1); + + StreamingResultSet delegate = + new ForwardingResultSet(mock(ResultSet.class)) { + private int currentRow = 0; + private Struct currentStruct; + + @Override + public boolean next() { + if (currentRow < totalRowCount) { + currentRow++; + currentStruct = Struct.newBuilder().set("ID").to((long) currentRow).build(); + return true; + } + return false; + } + + @Override + public Struct getCurrentRowAsStruct() { + return currentStruct; + } + + @Override + public boolean isDataAvailable() { + return currentRow < availableRows.get() || currentRow >= totalRowCount; + } + + @Override + public boolean initiateStreaming(AsyncResultSet.StreamMessageListener listener) { + registeredListener.set(listener); + streamingInitiated.countDown(); + return true; + } + }; + + Semaphore chunkRequested = new Semaphore(1); + AtomicInteger rowsReceived = new AtomicInteger(); + AtomicBoolean doneReceived = new AtomicBoolean(); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, bufferCapacity)) { + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + while (true) { + switch (resultSet.tryNext()) { + case OK: + int expectedId = rowsReceived.incrementAndGet(); + assertEquals(expectedId, resultSet.getLong("ID")); + break; + case NOT_READY: + if (availableRows.get() < totalRowCount + && chunkRequested.availablePermits() == 0) { + chunkRequested.release(); + } + return CallbackResponse.CONTINUE; + case DONE: + doneReceived.set(true); + return CallbackResponse.DONE; + } + } + }); + + assertTrue(streamingInitiated.await(5, TimeUnit.SECONDS)); + feederExecutor.execute( + () -> { + try { + while (availableRows.get() < totalRowCount) { + chunkRequested.acquire(); + availableRows.addAndGet(chunkSize); + AsyncResultSet.StreamMessageListener listener = registeredListener.get(); + if (listener != null) { + listener.onStreamMessage(PartialResultSet.getDefaultInstance(), false); + } + } + } catch (InterruptedException ignored) { + } + }); + + assertNull(future.get(30, TimeUnit.SECONDS)); + assertTrue(doneReceived.get()); + assertEquals(totalRowCount, rowsReceived.get()); + } finally { + feederExecutor.shutdownNow(); + callbackExecutor.shutdown(); + assertTrue(feederExecutor.awaitTermination(5, TimeUnit.SECONDS)); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testSetCallbackTwiceThrowsException() { + AsyncResultSetImpl resultSet = + new AsyncResultSetImpl( + mockedProvider, mock(ResultSet.class), AsyncResultSetImpl.DEFAULT_BUFFER_SIZE); + resultSet.setCallback(mock(Executor.class), mock(ReadyCallback.class)); + IllegalStateException exception = + assertThrows( + IllegalStateException.class, + () -> resultSet.setCallback(mock(Executor.class), mock(ReadyCallback.class))); + assertTrue( + "Expected message to contain 'callback may not be set multiple times', got: " + + exception.getMessage(), + exception.getMessage().contains("callback may not be set multiple times")); + } + + @Test + public void testCancelWithoutCallbackThrowsException() { + AsyncResultSetImpl resultSet = + new AsyncResultSetImpl( + mockedProvider, mock(ResultSet.class), AsyncResultSetImpl.DEFAULT_BUFFER_SIZE); + IllegalStateException exception = assertThrows(IllegalStateException.class, resultSet::cancel); + assertTrue( + "Expected message to contain 'cannot cancel a result set without a callback', got: " + + exception.getMessage(), + exception.getMessage().contains("cannot cancel a result set without a callback")); + } + + @Test + public void testResumeWithoutCallbackThrowsException() { + AsyncResultSetImpl resultSet = + new AsyncResultSetImpl( + mockedProvider, mock(ResultSet.class), AsyncResultSetImpl.DEFAULT_BUFFER_SIZE); + IllegalStateException exception = assertThrows(IllegalStateException.class, resultSet::resume); + assertTrue( + "Expected message to contain 'cannot resume a result set without a callback', got: " + + exception.getMessage(), + exception.getMessage().contains("cannot resume a result set without a callback")); + } + + @Test + public void testCancelAfterCompletionIsNoOp() throws Exception { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(false); + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl resultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + resultSet.setCallback( + callbackExecutor, + readyResultSet -> { + if (readyResultSet.tryNext() == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + future.get(5, TimeUnit.SECONDS); + assertEquals(AsyncResultSetImpl.State.DONE, resultSet.getState()); + resultSet.cancel(); + assertEquals(AsyncResultSetImpl.State.DONE, resultSet.getState()); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testResumeAfterCompletionIsNoOp() throws Exception { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(false); + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl resultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + resultSet.setCallback( + callbackExecutor, + readyResultSet -> { + if (readyResultSet.tryNext() == CursorState.DONE) { + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + future.get(5, TimeUnit.SECONDS); + assertEquals(AsyncResultSetImpl.State.DONE, resultSet.getState()); + resultSet.resume(); + assertEquals(AsyncResultSetImpl.State.DONE, resultSet.getState()); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testCloseDelegateResultSetIgnoresException() { + ResultSet delegate = mock(ResultSet.class); + doThrow(new RuntimeException("close failed")).when(delegate).close(); + AsyncResultSetImpl resultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE); + resultSet.close(); + assertTrue(resultSet.isClosed()); + verify(delegate).close(); + } + + @Test + public void testConcurrentResumeWhileConsumingWithPause() throws Exception { + int totalRowCount = 5; + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, true, true, true, true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + CountDownLatch callbackEnteredLatch = new CountDownLatch(1); + CountDownLatch resumeCalledLatch = new CountDownLatch(1); + AtomicInteger rowsReceived = new AtomicInteger(); + AtomicBoolean pausedOnce = new AtomicBoolean(); + + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + try (AsyncResultSetImpl resultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture future = + resultSet.setCallback( + callbackExecutor, + readyResultSet -> { + while (true) { + switch (readyResultSet.tryNext()) { + case OK: + int count = rowsReceived.incrementAndGet(); + if (count == 1 && pausedOnce.compareAndSet(false, true)) { + callbackEnteredLatch.countDown(); + try { + assertTrue(resumeCalledLatch.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException interruptedException) { + Thread.currentThread().interrupt(); + throw new RuntimeException(interruptedException); + } + return CallbackResponse.PAUSE; + } + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + + assertTrue(callbackEnteredLatch.await(5, TimeUnit.SECONDS)); + assertEquals(AsyncResultSetImpl.State.CONSUMING, resultSet.getState()); + resultSet.resume(); + resumeCalledLatch.countDown(); + + assertNull(future.get(5, TimeUnit.SECONDS)); + assertEquals(totalRowCount, rowsReceived.get()); + assertEquals(AsyncResultSetImpl.State.DONE, resultSet.getState()); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testVirtualThreadPerTaskExecutor() throws Exception { + Assume.assumeTrue(isJava21OrHigher()); + + ExecutorService virtualThreadExecutor = + ThreadFactoryUtil.tryCreateVirtualThreadPerTaskExecutor("async-result-set-virtual-thread"); + assertNotNull(virtualThreadExecutor); + try { + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()).thenReturn(true, true, true, false); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl( + simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + AtomicInteger rowCounter = new AtomicInteger(0); + AtomicBoolean ranOnVirtualThread = new AtomicBoolean(false); + ApiFuture future = + asyncResultSet.setCallback( + virtualThreadExecutor, + resultSet -> { + if (isVirtual(Thread.currentThread())) { + ranOnVirtualThread.set(true); + } + while (true) { + switch (resultSet.tryNext()) { + case OK: + rowCounter.incrementAndGet(); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + + future.get(5, TimeUnit.SECONDS); + assertEquals(3, rowCounter.get()); + assertTrue(ranOnVirtualThread.get()); + } + } finally { + virtualThreadExecutor.shutdown(); + assertTrue(virtualThreadExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + private static boolean isJava21OrHigher() { + String[] versionElements = System.getProperty("java.version").split("\\."); + int majorVersion = Integer.parseInt(versionElements[0]); + // Java 1.8 (Java 8) and lower used the format 1.8 etc. + // Java 9 and higher use the format 9.x + if (majorVersion == 1) { + majorVersion = Integer.parseInt(versionElements[1]); + } + return majorVersion >= 21; + } + + private static boolean isVirtual(Thread thread) { + if (thread == null) { + return false; + } + try { + java.lang.reflect.Method isVirtualMethod = Thread.class.getMethod("isVirtual"); + return (Boolean) isVirtualMethod.invoke(thread); + } catch (Exception ignore) { + return false; + } + } + + @Test + public void testCancelWhileProducerIsBlockedInNext() throws Exception { + CountDownLatch producerEnteredNext = new CountDownLatch(1); + CountDownLatch allowNextToComplete = new CountDownLatch(1); + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + + ResultSet delegate = mock(ResultSet.class); + when(delegate.next()) + .thenAnswer( + invocation -> { + producerEnteredNext.countDown(); + assertTrue(allowNextToComplete.await(5, TimeUnit.SECONDS)); + return true; + }); + when(delegate.getCurrentRowAsStruct()).thenReturn(mock(Struct.class)); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + ApiFuture callbackFuture = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + resultSet.tryNext(); + return CallbackResponse.CONTINUE; + }); + + // Wait until producer is actively executing delegate.next() + assertTrue(producerEnteredNext.await(5, TimeUnit.SECONDS)); + + // Cancel from another thread while producer is blocked + asyncResultSet.cancel(); + allowNextToComplete.countDown(); + + ExecutionException executionException = + assertThrows(ExecutionException.class, () -> callbackFuture.get(5, TimeUnit.SECONDS)); + assertTrue(executionException.getCause() instanceof SpannerException); + assertEquals( + ErrorCode.CANCELLED, ((SpannerException) executionException.getCause()).getErrorCode()); + verify(delegate, Mockito.atLeastOnce()).close(); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testSingleRowSplitAcrossChunksDeliveredCleanly() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + StreamingResultSet delegate = mock(StreamingResultSet.class); + + AtomicBoolean dataAvailable = new AtomicBoolean(false); + when(delegate.isDataAvailable()).thenAnswer(invocation -> dataAvailable.get()); + + when(delegate.next()).thenReturn(true, false); + Struct expectedRow = Struct.newBuilder().set("Value").to("chunked-payload").build(); + when(delegate.getCurrentRowAsStruct()).thenReturn(expectedRow); + + CountDownLatch initiateStreamingCalled = new CountDownLatch(1); + AtomicReference listenerRef = new AtomicReference<>(); + when(delegate.initiateStreaming(any(AsyncResultSet.StreamMessageListener.class))) + .thenAnswer( + invocation -> { + listenerRef.set(invocation.getArgument(0)); + initiateStreamingCalled.countDown(); + return true; + }); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + List receivedRows = new ArrayList<>(); + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + while (true) { + switch (resultSet.tryNext()) { + case OK: + receivedRows.add(resultSet.getCurrentRowAsStruct()); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + + assertTrue(initiateStreamingCalled.await(5, TimeUnit.SECONDS)); + assertNotNull(listenerRef.get()); + + // First chunk arrives with partial data: isDataAvailable is false. + // The producer yields without producing any row. + dataAvailable.set(false); + listenerRef + .get() + .onStreamMessage(PartialResultSet.newBuilder().setChunkedValue(true).build(), false); + + // Second chunk arrives completing the row: isDataAvailable is true. + dataAvailable.set(true); + listenerRef + .get() + .onStreamMessage(PartialResultSet.newBuilder().setChunkedValue(false).build(), false); + + future.get(5, TimeUnit.SECONDS); + assertEquals(1, receivedRows.size()); + assertEquals("chunked-payload", receivedRows.get(0).getString("Value")); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); + } + } + + @Test + public void testInitiateStreamingWithZeroRowsDeliversDoneImmediately() throws Exception { + ExecutorService callbackExecutor = Executors.newSingleThreadExecutor(); + StreamingResultSet delegate = mock(StreamingResultSet.class); + when(delegate.next()).thenReturn(false); + when(delegate.isDataAvailable()).thenReturn(true); + + CountDownLatch initiateStreamingCalled = new CountDownLatch(1); + AtomicReference listenerRef = new AtomicReference<>(); + when(delegate.initiateStreaming(any(AsyncResultSet.StreamMessageListener.class))) + .thenAnswer( + invocation -> { + listenerRef.set(invocation.getArgument(0)); + initiateStreamingCalled.countDown(); + return true; + }); + + try (AsyncResultSetImpl asyncResultSet = + new AsyncResultSetImpl(simpleProvider, delegate, AsyncResultSetImpl.DEFAULT_BUFFER_SIZE)) { + AtomicBoolean receivedDone = new AtomicBoolean(false); + ApiFuture future = + asyncResultSet.setCallback( + callbackExecutor, + resultSet -> { + CursorState state = resultSet.tryNext(); + if (state == CursorState.DONE) { + receivedDone.set(true); + return CallbackResponse.DONE; + } + return CallbackResponse.CONTINUE; + }); + + assertTrue(initiateStreamingCalled.await(5, TimeUnit.SECONDS)); + assertNotNull(listenerRef.get()); + + // Signal end of stream + listenerRef.get().onStreamMessage(GrpcStreamIterator.END_OF_STREAM, false); + + future.get(5, TimeUnit.SECONDS); + assertTrue(receivedDone.get()); + assertEquals(0, asyncResultSet.getBufferSize()); + } finally { + callbackExecutor.shutdown(); + assertTrue(callbackExecutor.awaitTermination(5, TimeUnit.SECONDS)); } } } diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ForwardingResultSetTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ForwardingResultSetTest.java new file mode 100644 index 000000000000..cb918487cbf7 --- /dev/null +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ForwardingResultSetTest.java @@ -0,0 +1,94 @@ +/* + * Copyright 2024 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.spanner; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import com.google.spanner.v1.ResultSetMetadata; +import com.google.spanner.v1.ResultSetStats; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public class ForwardingResultSetTest { + + private interface MockStreamingResultSet extends ResultSet, StreamingResultSet {} + + @Test + public void testDelegation() { + MockStreamingResultSet delegate = mock(MockStreamingResultSet.class); + when(delegate.next()).thenReturn(true, false); + ResultSetStats stats = ResultSetStats.getDefaultInstance(); + ResultSetMetadata metadata = ResultSetMetadata.getDefaultInstance(); + when(delegate.getStats()).thenReturn(stats); + when(delegate.getMetadata()).thenReturn(metadata); + + ForwardingResultSet forwardingResultSet = new ForwardingResultSet(delegate); + + assertTrue(forwardingResultSet.next()); + assertFalse(forwardingResultSet.next()); + assertEquals(stats, forwardingResultSet.getStats()); + assertEquals(metadata, forwardingResultSet.getMetadata()); + assertEquals(delegate, forwardingResultSet.getDelegate()); + + forwardingResultSet.close(); + verify(delegate).close(); + } + + @Test + public void testInitiateStreamingDelegation() { + MockStreamingResultSet delegate = mock(MockStreamingResultSet.class); + AsyncResultSet.StreamMessageListener listener = + mock(AsyncResultSet.StreamMessageListener.class); + when(delegate.initiateStreaming(listener)).thenReturn(true); + + ForwardingResultSet forwardingResultSet = new ForwardingResultSet(delegate); + assertTrue(forwardingResultSet.initiateStreaming(listener)); + verify(delegate).initiateStreaming(listener); + } + + @Test + public void testIsDataAvailableDelegation() { + MockStreamingResultSet delegate = mock(MockStreamingResultSet.class); + when(delegate.isDataAvailable()).thenReturn(false, true); + + ForwardingResultSet forwardingResultSet = new ForwardingResultSet(delegate); + assertFalse(forwardingResultSet.isDataAvailable()); + assertTrue(forwardingResultSet.isDataAvailable()); + } + + @Test + public void testReplaceDelegate() { + ResultSet firstDelegate = mock(ResultSet.class); + ResultSet secondDelegate = mock(ResultSet.class); + when(firstDelegate.next()).thenReturn(true); + when(secondDelegate.next()).thenReturn(false); + + ForwardingResultSet forwardingResultSet = new ForwardingResultSet(firstDelegate); + assertTrue(forwardingResultSet.next()); + + forwardingResultSet.replaceDelegate(secondDelegate); + assertFalse(forwardingResultSet.next()); + assertEquals(secondDelegate, forwardingResultSet.getDelegate()); + } +} diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcResultSetTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcResultSetTest.java index 4007c972c24e..26ea13da4dda 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcResultSetTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcResultSetTest.java @@ -22,6 +22,9 @@ import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertThrows; import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import com.google.api.gax.grpc.GrpcCallContext; import com.google.api.gax.rpc.ApiCallContext; @@ -1254,4 +1257,45 @@ public void shouldThrowDeadlineExceededIfLastTrueIsNotReceived() { assertEquals("DEADLINE_EXCEEDED: stream wait timeout", spannerException.getMessage()); consumer.onCompleted(); } + + @Test + public void testIsDataAvailable() { + assertFalse(resultSet.isDataAvailable()); + + consumer.onPartialResultSet( + PartialResultSet.newBuilder() + .setMetadata(makeMetadata(Type.struct(Type.StructField.of("f", Type.string())))) + .addValues(Value.string("val1").toProto()) + .addValues(Value.string("val2").toProto()) + .build()); + + assertTrue(resultSet.isDataAvailable()); + assertTrue(resultSet.next()); + assertEquals("val1", resultSet.getString(0)); + + assertTrue(resultSet.isDataAvailable()); + assertTrue(resultSet.next()); + assertEquals("val2", resultSet.getString(0)); + + assertFalse(resultSet.isDataAvailable()); + consumer.onCompleted(); + assertTrue(resultSet.isDataAvailable()); + assertFalse(resultSet.next()); + } + + @Test + public void testInitiateStreamingDelegation() { + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator mockStream = + mock(AbstractResultSet.CloseableIterator.class); + AsyncResultSet.StreamMessageListener listener = + mock(AsyncResultSet.StreamMessageListener.class); + when(mockStream.initiateStreaming(listener)).thenReturn(true); + + GrpcResultSet streamingResultSet = new GrpcResultSet(mockStream, new NoOpListener()); + assertTrue(streamingResultSet.initiateStreaming(listener)); + verify(mockStream).initiateStreaming(listener); + + assertFalse(resultSet.initiateStreaming(listener)); + } } diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcValueIteratorTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcValueIteratorTest.java new file mode 100644 index 000000000000..aeb31b8fe4af --- /dev/null +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/GrpcValueIteratorTest.java @@ -0,0 +1,96 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.spanner; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import com.google.cloud.spanner.AbstractResultSet.CloseableIterator; +import com.google.protobuf.Value; +import com.google.spanner.v1.PartialResultSet; +import com.google.spanner.v1.ResultSetMetadata; +import com.google.spanner.v1.StructType; +import com.google.spanner.v1.StructType.Field; +import com.google.spanner.v1.Type; +import com.google.spanner.v1.TypeCode; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public class GrpcValueIteratorTest { + @SuppressWarnings("unchecked") + private final CloseableIterator stream = mock(CloseableIterator.class); + + private final AbstractResultSet.Listener listener = mock(AbstractResultSet.Listener.class); + + @Test + public void testIsDataAvailableWhenCurrentIsNull() { + GrpcValueIterator iterator = new GrpcValueIterator(stream, listener); + when(stream.isDataAvailable()).thenReturn(false); + assertFalse(iterator.isDataAvailable()); + + when(stream.isDataAvailable()).thenReturn(true); + assertTrue(iterator.isDataAvailable()); + } + + @Test + public void testIsDataAvailableWhenValuesRemainInCurrent() { + PartialResultSet prs = + PartialResultSet.newBuilder() + .setMetadata( + ResultSetMetadata.newBuilder() + .setRowType( + StructType.newBuilder() + .addFields( + Field.newBuilder() + .setName("col") + .setType(Type.newBuilder().setCode(TypeCode.STRING))))) + .addValues(Value.newBuilder().setStringValue("val1")) + .addValues(Value.newBuilder().setStringValue("val2")) + .build(); + when(stream.hasNext()).thenReturn(true, false); + when(stream.next()).thenReturn(prs); + + GrpcValueIterator iterator = new GrpcValueIterator(stream, listener); + when(stream.isDataAvailable()).thenReturn(false); + + // Initial state before reading: current is null, stream has no data available + assertFalse(iterator.isDataAvailable()); + + // Advance iterator: reads first value "val1", current is set, pos is now 1 (1 < 2 values) + assertTrue(iterator.hasNext()); + assertEquals("val1", iterator.next().getStringValue()); + + // Even though stream.isDataAvailable() is false, iterator has 1 value remaining in current + // chunk + assertFalse(stream.isDataAvailable()); + assertTrue(iterator.isDataAvailable()); + + // Read second value "val2", pos is now 2 (all 2 values consumed from current chunk) + assertTrue(iterator.hasNext()); + assertEquals("val2", iterator.next().getStringValue()); + + // Now all values in current chunk are consumed, so isDataAvailable matches stream + assertFalse(iterator.isDataAvailable()); + when(stream.isDataAvailable()).thenReturn(true); + assertTrue(iterator.isDataAvailable()); + } +} diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ReadAsyncTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ReadAsyncTest.java index 1251ee270faf..e3d075b75cb9 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ReadAsyncTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ReadAsyncTest.java @@ -19,7 +19,10 @@ import static com.google.cloud.spanner.MockSpannerTestUtil.*; import static com.google.cloud.spanner.SpannerApiFutures.get; import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNull; import static org.junit.Assert.assertThrows; +import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import com.google.api.core.ApiFuture; @@ -33,6 +36,9 @@ import com.google.common.collect.ContiguousSet; import com.google.common.collect.ImmutableSet; import com.google.common.collect.Iterables; +import com.google.protobuf.ByteString; +import com.google.spanner.v1.ExecuteSqlRequest; +import com.google.spanner.v1.ReadRequest; import io.grpc.Server; import io.grpc.Status; import io.grpc.inprocess.InProcessServerBuilder; @@ -49,6 +55,7 @@ import java.util.concurrent.Executors; import java.util.concurrent.ScheduledThreadPoolExecutor; import java.util.concurrent.SynchronousQueue; +import java.util.concurrent.TimeUnit; import org.junit.After; import org.junit.AfterClass; import org.junit.Before; @@ -114,6 +121,7 @@ public void before() { public void after() { spanner.close(); mockSpanner.removeAllExecutionTimes(); + mockSpanner.clearRequests(); } @Test @@ -431,6 +439,115 @@ public void cancel() throws Exception { assertThat(values).containsExactly("v1"); } + @Test + public void readAsyncRetriesOnUnavailableHalfway() throws Exception { + int totalRowCount = 50; + int errorIndex = 20; + String retryTableName = "RetryTable"; + mockSpanner.putStatementResult( + StatementResult.read( + retryTableName, + KeySet.all(), + READ_COLUMN_NAMES, + generateKeyValueResultSet(ContiguousSet.closed(1, totalRowCount)))); + mockSpanner.setStreamingReadExecutionTime( + SimulatedExecutionTime.ofStreamException( + Status.UNAVAILABLE.asRuntimeException(), errorIndex)); + mockSpanner.clearRequests(); + + List receivedKeys = new ArrayList<>(); + List receivedValues = new ArrayList<>(); + try (AsyncResultSet resultSet = + client.singleUse().readAsync(retryTableName, KeySet.all(), READ_COLUMN_NAMES)) { + ApiFuture future = + resultSet.setCallback( + executor, + ready -> { + while (true) { + switch (ready.tryNext()) { + case OK: + receivedKeys.add(ready.getString("Key")); + receivedValues.add(ready.getString("Value")); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + assertNull(future.get(10, TimeUnit.SECONDS)); + } + + assertEquals(totalRowCount, receivedKeys.size()); + assertEquals(totalRowCount, receivedValues.size()); + for (int i = 0; i < totalRowCount; i++) { + assertEquals("k" + (i + 1), receivedKeys.get(i)); + assertEquals("v" + (i + 1), receivedValues.get(i)); + } + + assertEquals(2, mockSpanner.countRequestsOfType(ReadRequest.class)); + ReadRequest initialRequest = mockSpanner.getRequestsOfType(ReadRequest.class).get(0); + assertTrue(initialRequest.getResumeToken().isEmpty()); + + ReadRequest resumeRequest = mockSpanner.getRequestsOfType(ReadRequest.class).get(1); + assertEquals( + ByteString.copyFromUtf8(String.format("%09d", errorIndex)), resumeRequest.getResumeToken()); + } + + @Test + public void executeQueryAsyncRetriesOnUnavailableHalfway() throws Exception { + int totalRowCount = 50; + int errorIndex = 20; + Statement statement = Statement.of("SELECT Key, Value FROM RetryTable"); + mockSpanner.putStatementResult( + StatementResult.query( + statement, generateKeyValueResultSet(ContiguousSet.closed(1, totalRowCount)))); + mockSpanner.setExecuteStreamingSqlExecutionTime( + SimulatedExecutionTime.ofStreamException( + Status.UNAVAILABLE.asRuntimeException(), errorIndex)); + mockSpanner.clearRequests(); + + List receivedKeys = new ArrayList<>(); + List receivedValues = new ArrayList<>(); + try (AsyncResultSet resultSet = client.singleUse().executeQueryAsync(statement)) { + ApiFuture future = + resultSet.setCallback( + executor, + ready -> { + while (true) { + switch (ready.tryNext()) { + case OK: + receivedKeys.add(ready.getString("Key")); + receivedValues.add(ready.getString("Value")); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + assertNull(future.get(10, TimeUnit.SECONDS)); + } + + assertEquals(totalRowCount, receivedKeys.size()); + assertEquals(totalRowCount, receivedValues.size()); + for (int i = 0; i < totalRowCount; i++) { + assertEquals("k" + (i + 1), receivedKeys.get(i)); + assertEquals("v" + (i + 1), receivedValues.get(i)); + } + + assertEquals(2, mockSpanner.countRequestsOfType(ExecuteSqlRequest.class)); + ExecuteSqlRequest initialRequest = + mockSpanner.getRequestsOfType(ExecuteSqlRequest.class).get(0); + assertTrue(initialRequest.getResumeToken().isEmpty()); + + ExecuteSqlRequest resumeRequest = mockSpanner.getRequestsOfType(ExecuteSqlRequest.class).get(1); + assertEquals( + ByteString.copyFromUtf8(String.format("%09d", errorIndex)), resumeRequest.getResumeToken()); + } + private boolean isMultiplexedSessionsEnabled() { if (spanner.getOptions() == null || spanner.getOptions().getSessionPoolOptions() == null) { return false; diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ResumableStreamIteratorTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ResumableStreamIteratorTest.java index f13c0bb1237a..90410d13be02 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ResumableStreamIteratorTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/ResumableStreamIteratorTest.java @@ -18,23 +18,29 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertThrows; +import static org.junit.Assert.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; import com.google.api.client.util.BackOff; +import com.google.api.client.util.ExponentialBackOff; +import com.google.api.gax.retrying.RetrySettings; +import com.google.api.gax.rpc.StatusCode.Code; import com.google.cloud.spanner.ErrorHandler.DefaultErrorHandler; import com.google.cloud.spanner.XGoogSpannerRequestId.NoopRequestIdCreator; +import com.google.cloud.spanner.spi.v1.SpannerRpc; import com.google.cloud.spanner.v1.stub.SpannerStubSettings; import com.google.common.collect.AbstractIterator; import com.google.common.collect.ImmutableList; import com.google.common.collect.Lists; import com.google.protobuf.ByteString; -import com.google.protobuf.Duration; import com.google.protobuf.Value; import com.google.rpc.RetryInfo; import com.google.spanner.v1.PartialResultSet; +import io.grpc.Context; import io.grpc.Metadata; import io.grpc.Status; import io.grpc.StatusRuntimeException; @@ -46,10 +52,14 @@ import java.io.IOException; import java.lang.reflect.Field; import java.util.ArrayList; +import java.util.Collections; import java.util.Iterator; import java.util.LinkedList; import java.util.List; +import java.util.NoSuchElementException; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import javax.annotation.Nullable; import org.junit.Assume; import org.junit.Before; @@ -59,6 +69,7 @@ import org.junit.runners.Parameterized.Parameter; import org.junit.runners.Parameterized.Parameters; import org.mockito.Mockito; +import org.threeten.bp.Duration; /** Unit tests for {@link ResumableStreamIterator}. */ @RunWith(Parameterized.class) @@ -89,7 +100,7 @@ private static StatusRuntimeException statusWithRetryInfo(ErrorCode code) { RetryInfo retryInfo = RetryInfo.newBuilder() .setRetryDelay( - Duration.newBuilder() + com.google.protobuf.Duration.newBuilder() .setNanos((int) TimeUnit.MILLISECONDS.toNanos(1L)) .setSeconds(0L)) .build(); @@ -175,7 +186,7 @@ AbstractResultSet.CloseableIterator startStream( @Nullable ByteString resumeToken, AsyncResultSet.StreamMessageListener streamMessageListener, XGoogSpannerRequestId requestId) { - return starter.startStream(resumeToken, null); + return starter.startStream(resumeToken, streamMessageListener); } }; } @@ -472,6 +483,645 @@ public void bufferLimitMissingTokensSafeToRetry() { assertThat(consume(resumableStreamIterator)).containsExactly("a", "b", "c", "d").inOrder(); } + @Test + public void isDataAvailableWithGrpcStreamIterator() { + initWithLimit(512); + GrpcStreamIterator grpcStream = new GrpcStreamIterator(false, 4, false); + Mockito.when(starter.startStream(null, null)).thenReturn(grpcStream); + resumableStreamIterator.setStream(grpcStream); + + // Empty stream -> false + assertFalse(resumableStreamIterator.isDataAvailable()); + + // Chunk without resume token arrives -> immediately available because GrpcStreamIterator has + // data + grpcStream.consumer().onPartialResultSet(resultSet(null, "a")); + assertTrue(resumableStreamIterator.isDataAvailable()); + } + + @Test + public void isDataAvailableWithoutResumeTokensWithProductionBufferSize() throws Exception { + initWithLimit(512); + GrpcStreamIterator grpcStream = new GrpcStreamIterator(false, 4, false); + SpannerRpc.StreamingCall call = Mockito.mock(SpannerRpc.StreamingCall.class); + grpcStream.setCall(call, false); + Mockito.when(starter.startStream(null, null)).thenReturn(grpcStream); + resumableStreamIterator.setStream(grpcStream); + + assertFalse(resumableStreamIterator.isDataAvailable()); + + java.util.concurrent.ExecutorService feeder = + java.util.concurrent.Executors.newSingleThreadExecutor(); + try { + feeder.submit( + () -> { + for (int i = 0; i < 10; i++) { + grpcStream.consumer().onPartialResultSet(resultSet(null, "val" + i)); + } + grpcStream.consumer().onCompleted(); + }); + + List results = new ArrayList<>(); + while (resumableStreamIterator.hasNext()) { + PartialResultSet prs = resumableStreamIterator.next(); + results.add(prs.getValues(0).getStringValue()); + } + assertEquals(10, results.size()); + for (int i = 0; i < 10; i++) { + assertEquals("val" + i, results.get(i)); + } + } finally { + feeder.shutdown(); + } + } + + @Test + public void isDataAvailableWithResumeTokenInGrpcStream() { + initWithLimit(10); + GrpcStreamIterator grpcStream = new GrpcStreamIterator(false, 4, false); + Mockito.when(starter.startStream(null, null)).thenReturn(grpcStream); + resumableStreamIterator.setStream(grpcStream); + + assertFalse(resumableStreamIterator.isDataAvailable()); + + // Chunk with resume token arrives -> immediately ready to emit! + grpcStream.consumer().onPartialResultSet(resultSet(ByteString.copyFromUtf8("r1"), "a")); + assertTrue(resumableStreamIterator.isDataAvailable()); + } + + @Test + public void isDataAvailableWhenStreamIsNull() { + initWithLimit(10); + resumableStreamIterator.setStream(null); + assertFalse(resumableStreamIterator.isDataAvailable()); + } + + @Test + public void isDataAvailableWhenFinished() { + initWithLimit(10); + setInternalState(ResumableStreamIterator.class, resumableStreamIterator, "finished", true); + assertTrue(resumableStreamIterator.isDataAvailable()); + } + + @Test + public void isDataAvailableWhenClosedReturnsTrue() { + initWithLimit(10); + resumableStreamIterator.close("closed"); + assertTrue(resumableStreamIterator.isClosed()); + assertTrue(resumableStreamIterator.isDataAvailable()); + } + + @Test + public void closeDuringStartStreamClosesStreamAndPreventsResurrection() { + initWithLimit(10); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + when(starter.startStream(null, null)) + .thenAnswer( + invocation -> { + resumableStreamIterator.close("cancelled while starting stream"); + return streamIterator; + }); + + assertFalse(resumableStreamIterator.hasNext()); + assertTrue(resumableStreamIterator.isClosed()); + verify(streamIterator).close(null); + } + + @Test + public void setStreamWhenAlreadyClosedImmediatelyClosesStream() { + initWithLimit(10); + resumableStreamIterator.close("already closed"); + assertTrue(resumableStreamIterator.isClosed()); + + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + resumableStreamIterator.setStream(streamIterator); + + verify(streamIterator).close(null); + } + + @Test + public void multipleCloseCallsAreIdempotent() { + initWithLimit(10); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + resumableStreamIterator.setStream(streamIterator); + + resumableStreamIterator.close("first close"); + resumableStreamIterator.close("second close"); + + assertTrue(resumableStreamIterator.isClosed()); + verify(streamIterator, Mockito.times(1)).close("first close"); + } + + @Test + public void closeWhenStreamIsNullEndsSpan() { + SpannerOptions.resetActiveTracingFramework(); + SpannerOptions.enableOpenTelemetryTraces(); + + io.opentelemetry.api.trace.Span openTelemetrySpan = mock(io.opentelemetry.api.trace.Span.class); + ISpan span = new OpenTelemetrySpan(openTelemetrySpan); + TraceWrapper tracer = mock(TraceWrapper.class); + when(tracer.spanBuilderWithExplicitParent(Mockito.anyString(), Mockito.any(), Mockito.any())) + .thenReturn(span); + + ResumableStreamIterator iterator = + new ResumableStreamIterator( + 10, + "testStream", + span, + tracer, + DefaultErrorHandler.INSTANCE, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + }; + + iterator.close("cancelled before streaming"); + assertTrue(iterator.isClosed()); + verify(openTelemetrySpan).end(); + } + + @Test + public void closeWithNullMessage() { + initWithLimit(10); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + resumableStreamIterator.setStream(streamIterator); + + resumableStreamIterator.close(null); + assertTrue(resumableStreamIterator.isClosed()); + verify(streamIterator).close(null); + } + + @Test + public void isDataAvailableWhenBufferHasItems() { + initWithLimit(10); + ResultSetStream s1 = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(s1)); + when(s1.next()) + .thenReturn(resultSet(null, "chunkWithoutToken")) + .thenReturn(resultSet(ByteString.copyFromUtf8("r1"), "chunkWithToken")) + .thenReturn(null); + + // Consuming the first element causes computeNext to buffer both chunks (up to the resume token) + // and pop the first one, leaving chunkWithToken in the buffer. + assertEquals("chunkWithoutToken", resumableStreamIterator.next().getValues(0).getStringValue()); + + // With chunkWithToken still buffered, isDataAvailable() returns true via buffer.isEmpty() == + // false + assertTrue(resumableStreamIterator.isDataAvailable()); + assertEquals("chunkWithToken", resumableStreamIterator.next().getValues(0).getStringValue()); + } + + @Test + public void isWithBeginTransactionDelegatesToStream() { + initWithLimit(10); + assertFalse(resumableStreamIterator.isWithBeginTransaction()); + + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + when(streamIterator.isWithBeginTransaction()).thenReturn(true); + resumableStreamIterator.setStream(streamIterator); + assertTrue(resumableStreamIterator.isWithBeginTransaction()); + + when(streamIterator.isWithBeginTransaction()).thenReturn(false); + assertFalse(resumableStreamIterator.isWithBeginTransaction()); + } + + @Test + public void isLastStatementDelegatesToStream() { + initWithLimit(10); + assertFalse(resumableStreamIterator.isLastStatement()); + + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + when(streamIterator.isLastStatement()).thenReturn(true); + resumableStreamIterator.setStream(streamIterator); + assertTrue(resumableStreamIterator.isLastStatement()); + + when(streamIterator.isLastStatement()).thenReturn(false); + assertFalse(resumableStreamIterator.isLastStatement()); + } + + @Test + public void initiateStreamingStartsStreamAndRegistersListener() { + initWithLimit(10); + AsyncResultSet.StreamMessageListener streamMessageListener = + mock(AsyncResultSet.StreamMessageListener.class); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + when(starter.startStream(null, streamMessageListener)).thenReturn(streamIterator); + + assertTrue(resumableStreamIterator.initiateStreaming(streamMessageListener)); + verify(starter).startStream(null, streamMessageListener); + verify(streamIterator).requestPrefetchChunks(); + } + + @Test + public void setStreamDoesNotRequestPrefetchChunks() { + initWithLimit(10); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + resumableStreamIterator.setStream(streamIterator); + verify(streamIterator, Mockito.never()).requestPrefetchChunks(); + } + + @Test + public void computeNextWhenAlreadyClosedReturnsEndOfData() { + initWithLimit(10); + resumableStreamIterator.close("pre-closed"); + assertFalse(resumableStreamIterator.hasNext()); + assertThrows(NoSuchElementException.class, () -> resumableStreamIterator.next()); + } + + @Test + public void retryOnDifferentGrpcChannelWhenSupported() { + ErrorHandler errorHandler = mock(ErrorHandler.class); + when(errorHandler.getMaxAttempts()).thenReturn(2); + when(errorHandler.translateException(Mockito.any())) + .thenReturn(new RetryOnDifferentGrpcChannelException("channel failed", 1, null)); + + final AtomicInteger channelRetries = new AtomicInteger(0); + ResumableStreamIterator channelRetryIterator = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + errorHandler, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + + @Override + boolean prepareIteratorForRetryOnDifferentGrpcChannel() { + channelRetries.incrementAndGet(); + return true; + } + }; + + ResultSetStream firstStream = mock(ResultSetStream.class); + ResultSetStream secondStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)) + .thenReturn(new ResultSetIterator(firstStream)) + .thenReturn(new ResultSetIterator(secondStream)); + + when(firstStream.next()) + .thenThrow(new NonRetryableException(ErrorCode.UNAVAILABLE, "transient channel error")); + when(secondStream.next()) + .thenReturn(resultSet(ByteString.copyFromUtf8("r1"), "recoveredValue")) + .thenReturn(null); + + List results = consume(channelRetryIterator); + assertEquals(1, results.size()); + assertEquals("recoveredValue", results.get(0)); + assertEquals(1, channelRetries.get()); + } + + @Test + public void retryOnDifferentGrpcChannelNotAttemptedWhenPrepareReturnsFalse() { + ErrorHandler errorHandler = mock(ErrorHandler.class); + when(errorHandler.getMaxAttempts()).thenReturn(2); + when(errorHandler.translateException(Mockito.any())) + .thenReturn(new RetryOnDifferentGrpcChannelException("channel failed", 1, null)); + + ResumableStreamIterator channelRetryIterator = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + errorHandler, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + }; + + ResultSetStream firstStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(firstStream)); + when(firstStream.next()) + .thenThrow(new NonRetryableException(ErrorCode.UNAVAILABLE, "transient error")); + + SpannerException spannerException = + assertThrows(SpannerException.class, channelRetryIterator::next); + assertEquals(ErrorCode.UNAVAILABLE, spannerException.getErrorCode()); + } + + @Test + public void retryOnDifferentGrpcChannelNotAttemptedWhenMaxAttemptsExceeded() { + ErrorHandler errorHandler = mock(ErrorHandler.class); + when(errorHandler.getMaxAttempts()).thenReturn(1); + when(errorHandler.translateException(Mockito.any())) + .thenReturn(new RetryOnDifferentGrpcChannelException("channel failed", 1, null)); + + ResumableStreamIterator channelRetryIterator = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + errorHandler, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + + @Override + boolean prepareIteratorForRetryOnDifferentGrpcChannel() { + return true; + } + }; + + ResultSetStream firstStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(firstStream)); + when(firstStream.next()) + .thenThrow(new NonRetryableException(ErrorCode.UNAVAILABLE, "transient error")); + + SpannerException spannerException = + assertThrows(SpannerException.class, channelRetryIterator::next); + assertEquals(ErrorCode.UNAVAILABLE, spannerException.getErrorCode()); + } + + @Test + public void retryOnDifferentGrpcChannelNotAttemptedWhenResumeTokenPresent() { + ErrorHandler errorHandler = mock(ErrorHandler.class); + when(errorHandler.getMaxAttempts()).thenReturn(5); + when(errorHandler.translateException(Mockito.any())) + .thenReturn(new RetryOnDifferentGrpcChannelException("channel failed", 1, null)); + + final AtomicBoolean prepareCalled = new AtomicBoolean(false); + ResumableStreamIterator channelRetryIterator = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + errorHandler, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + + @Override + boolean prepareIteratorForRetryOnDifferentGrpcChannel() { + prepareCalled.set(true); + return true; + } + }; + + ResultSetStream firstStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(firstStream)); + when(firstStream.next()) + .thenReturn(resultSet(ByteString.copyFromUtf8("r1"), "firstChunk")) + .thenThrow(new NonRetryableException(ErrorCode.UNAVAILABLE, "error after token")); + + assertEquals("firstChunk", channelRetryIterator.next().getValues(0).getStringValue()); + SpannerException spannerException = + assertThrows(SpannerException.class, channelRetryIterator::next); + assertEquals(ErrorCode.UNAVAILABLE, spannerException.getErrorCode()); + assertFalse(prepareCalled.get()); + } + + @Test + public void newBackOffWithDefaultSettings() { + initWithLimit(10); + ExponentialBackOff backOff = resumableStreamIterator.newBackOff(); + assertEquals(1.0, backOff.getMultiplier(), 0.001); + assertEquals(10, backOff.getInitialIntervalMillis()); + assertEquals(1000, backOff.getMaxIntervalMillis()); + assertEquals(Integer.MAX_VALUE, backOff.getMaxElapsedTimeMillis()); + } + + @Test + public void newBackOffWithCustomSettings() { + RetrySettings customRetrySettings = + RetrySettings.newBuilder() + .setInitialRetryDelay(Duration.ofMillis(50)) + .setMaxRetryDelay(Duration.ofMillis(500)) + .setRetryDelayMultiplier(2.0) + .setTotalTimeout(Duration.ofMillis(5000)) + .build(); + ResumableStreamIterator customIterator = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + DefaultErrorHandler.INSTANCE, + customRetrySettings, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return starter.startStream(resumeToken, streamMessageListener); + } + }; + + ExponentialBackOff backOff = customIterator.newBackOff(); + assertEquals(2.0, backOff.getMultiplier(), 0.001); + assertEquals(50, backOff.getInitialIntervalMillis()); + assertEquals(500, backOff.getMaxIntervalMillis()); + assertEquals(5000, backOff.getMaxElapsedTimeMillis()); + } + + @Test + public void nextBackOffMillisWhenIOExceptionThrowsSpannerException() throws Exception { + BackOff mockBackOff = mock(BackOff.class); + when(mockBackOff.nextBackOffMillis()).thenThrow(new IOException("Simulated I/O error")); + + SpannerException spannerException = + assertThrows( + SpannerException.class, () -> ResumableStreamIterator.nextBackOffMillis(mockBackOff)); + assertEquals(ErrorCode.INTERNAL, spannerException.getErrorCode()); + } + + @Test + public void backoffSleepWhenContextCancelledThrowsCancellationException() { + initWithLimit(10); + ResultSetStream firstStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(firstStream)); + when(firstStream.next()) + .thenThrow(new RetryableException(errorCodeParameter, "transient error")); + + Context.CancellableContext cancellableContext = Context.current().withCancellation(); + cancellableContext.cancel(new RuntimeException("context cancelled by test")); + + cancellableContext.run( + () -> { + SpannerException spannerException = + assertThrows(SpannerException.class, resumableStreamIterator::next); + assertEquals(ErrorCode.CANCELLED, spannerException.getErrorCode()); + }); + } + + @Test + public void unhandledRuntimeExceptionSetsSpanStatusAndThrows() { + initWithLimit(10); + ResultSetStream firstStream = mock(ResultSetStream.class); + when(starter.startStream(null, null)).thenReturn(new ResultSetIterator(firstStream)); + when(firstStream.next()).thenThrow(new IllegalStateException("unexpected runtime exception")); + + IllegalStateException thrown = + assertThrows(IllegalStateException.class, resumableStreamIterator::next); + assertEquals("unexpected runtime exception", thrown.getMessage()); + } + + @Test + public void isRetryableMatchesRetryableCodesEvenIfExceptionIsNotRetryable() { + ResumableStreamIterator iteratorWithRetryableCodes = + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + DefaultErrorHandler.INSTANCE, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + Collections.singleton(Code.UNAVAILABLE), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return null; + } + }; + + NonRetryableException exceptionWithUnavailableCode = + new NonRetryableException(ErrorCode.UNAVAILABLE, "non retryable flag but UNAVAILABLE code"); + assertTrue(iteratorWithRetryableCodes.isRetryable(exceptionWithUnavailableCode)); + + NonRetryableException exceptionWithInvalidArgument = + new NonRetryableException(ErrorCode.INVALID_ARGUMENT, "non retryable code"); + assertFalse(iteratorWithRetryableCodes.isRetryable(exceptionWithInvalidArgument)); + } + + @Test + public void constructorNegativeMaxBufferSizeThrowsIllegalArgumentException() { + assertThrows(IllegalArgumentException.class, () -> initWithLimit(-1)); + } + + @Test + public void constructorNullRetrySettingsThrowsNullPointerException() { + assertThrows( + NullPointerException.class, + () -> + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + DefaultErrorHandler.INSTANCE, + null, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetryableCodes(), + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return null; + } + }); + } + + @Test + public void constructorNullRetryableCodesThrowsNullPointerException() { + assertThrows( + NullPointerException.class, + () -> + new ResumableStreamIterator( + 10, + "", + new OpenTelemetrySpan(mock(io.opentelemetry.api.trace.Span.class)), + new TraceWrapper(Tracing.getTracer(), OpenTelemetry.noop().getTracer(""), false), + DefaultErrorHandler.INSTANCE, + SpannerStubSettings.newBuilder().executeStreamingSqlSettings().getRetrySettings(), + null, + NoopRequestIdCreator.INSTANCE) { + @Override + AbstractResultSet.CloseableIterator startStream( + @Nullable ByteString resumeToken, + AsyncResultSet.StreamMessageListener streamMessageListener, + XGoogSpannerRequestId requestId) { + return null; + } + }); + } + + @Test + public void isDataAvailableWhenBufferHasItemsWithoutResumeTokenDelegatesToUnderlyingStream() { + initWithLimit(10); + @SuppressWarnings("unchecked") + AbstractResultSet.CloseableIterator streamIterator = + mock(AbstractResultSet.CloseableIterator.class); + resumableStreamIterator.setStream(streamIterator); + + PartialResultSet chunkWithoutToken = resultSet(null, "chunk1"); + synchronized (resumableStreamIterator.buffer) { + resumableStreamIterator.buffer.add(chunkWithoutToken); + } + + when(streamIterator.isDataAvailable()).thenReturn(false); + assertFalse(resumableStreamIterator.isDataAvailable()); + + when(streamIterator.isDataAvailable()).thenReturn(true); + assertTrue(resumableStreamIterator.isDataAvailable()); + + PartialResultSet chunkWithToken = resultSet(ByteString.copyFromUtf8("token1"), "chunk2"); + synchronized (resumableStreamIterator.buffer) { + resumableStreamIterator.buffer.add(chunkWithToken); + } + when(streamIterator.isDataAvailable()).thenReturn(false); + assertTrue(resumableStreamIterator.isDataAvailable()); + } + static PartialResultSet resultSet(@Nullable ByteString resumeToken, String... data) { PartialResultSet.Builder builder = PartialResultSet.newBuilder(); if (resumeToken != null) { diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/StreamingUtilTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/StreamingUtilTest.java new file mode 100644 index 000000000000..1b61cce8683c --- /dev/null +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/StreamingUtilTest.java @@ -0,0 +1,93 @@ +/* + * Copyright 2024 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.cloud.spanner; + +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public class StreamingUtilTest { + + private interface MockStreamingResultSet extends ResultSet, StreamingResultSet {} + + @Test + public void testInitiateStreamingOnStreamingResultSet() { + MockStreamingResultSet streamingResultSet = mock(MockStreamingResultSet.class); + AsyncResultSet.StreamMessageListener listener = + mock(AsyncResultSet.StreamMessageListener.class); + when(streamingResultSet.initiateStreaming(listener)).thenReturn(true); + + assertTrue(StreamingUtil.initiateStreaming(streamingResultSet, listener)); + verify(streamingResultSet).initiateStreaming(listener); + } + + @Test + public void testInitiateStreamingOnNonStreamingResultSetReturnsFalse() { + ResultSet nonStreamingResultSet = mock(ResultSet.class); + AsyncResultSet.StreamMessageListener listener = + mock(AsyncResultSet.StreamMessageListener.class); + + assertFalse(StreamingUtil.initiateStreaming(nonStreamingResultSet, listener)); + } + + @Test + public void testIsDataAvailableOnStreamingResultSet() { + MockStreamingResultSet streamingResultSet = mock(MockStreamingResultSet.class); + when(streamingResultSet.isDataAvailable()).thenReturn(false, true); + + assertFalse(StreamingUtil.isDataAvailable(streamingResultSet)); + assertTrue(StreamingUtil.isDataAvailable(streamingResultSet)); + } + + @Test + public void testIsDataAvailableOnNonStreamingResultSetReturnsTrue() { + ResultSet nonStreamingResultSet = mock(ResultSet.class); + assertTrue(StreamingUtil.isDataAvailable(nonStreamingResultSet)); + } + + @Test + public void testIsDataAvailableOnNestedForwardingResultSet() { + MockStreamingResultSet underlyingStreamingResultSet = mock(MockStreamingResultSet.class); + ForwardingResultSet innerForwarder = new ForwardingResultSet(underlyingStreamingResultSet); + ForwardingResultSet outerForwarder = new ForwardingResultSet(innerForwarder); + + when(underlyingStreamingResultSet.isDataAvailable()).thenReturn(false); + assertFalse(StreamingUtil.isDataAvailable(outerForwarder)); + + when(underlyingStreamingResultSet.isDataAvailable()).thenReturn(true); + assertTrue(StreamingUtil.isDataAvailable(outerForwarder)); + } + + @Test + public void testInitiateStreamingOnForwardingResultSet() { + MockStreamingResultSet underlyingStreamingResultSet = mock(MockStreamingResultSet.class); + ForwardingResultSet forwarder = new ForwardingResultSet(underlyingStreamingResultSet); + AsyncResultSet.StreamMessageListener listener = + mock(AsyncResultSet.StreamMessageListener.class); + + when(underlyingStreamingResultSet.initiateStreaming(listener)).thenReturn(true); + assertTrue(StreamingUtil.initiateStreaming(forwarder, listener)); + verify(underlyingStreamingResultSet).initiateStreaming(listener); + } +} diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/it/ITLargeReadTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/it/ITLargeReadTest.java index 83d505e21241..fd594567bcd3 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/it/ITLargeReadTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/it/ITLargeReadTest.java @@ -17,8 +17,13 @@ package com.google.cloud.spanner.it; import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; +import com.google.api.core.ApiFuture; import com.google.cloud.ByteArray; +import com.google.cloud.spanner.AsyncResultSet; +import com.google.cloud.spanner.AsyncResultSet.CallbackResponse; import com.google.cloud.spanner.Database; import com.google.cloud.spanner.DatabaseClient; import com.google.cloud.spanner.Dialect; @@ -38,6 +43,10 @@ import java.util.Collections; import java.util.List; import java.util.Random; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.AfterClass; import org.junit.BeforeClass; import org.junit.ClassRule; @@ -217,4 +226,87 @@ private void validate(ResultSet resultSet) { } assertThat(i).isEqualTo(numRows); } + + @Test + public void readAsync() throws Exception { + try (AsyncResultSet resultSet = + getClient(dialect.dialect) + .singleUse() + .readAsync( + TABLE_NAME, KeySet.all(), Arrays.asList("Key", "Data", "Fingerprint", "Size"))) { + validateAsync(resultSet); + } + } + + @Test + public void readAsyncWithSmallPrefetchChunks() throws Exception { + try (AsyncResultSet resultSet = + getClient(dialect.dialect) + .singleUse() + .readAsync( + TABLE_NAME, + KeySet.all(), + Arrays.asList("Key", "Data", "Fingerprint", "Size"), + Options.prefetchChunks(1))) { + validateAsync(resultSet); + } + } + + @Test + public void queryAsync() throws Exception { + try (AsyncResultSet resultSet = + getClient(dialect.dialect) + .singleUse() + .executeQueryAsync( + Statement.of( + "SELECT Key, Data, Fingerprint, Size FROM " + TABLE_NAME + " ORDER BY Key"))) { + validateAsync(resultSet); + } + } + + @Test + public void queryAsyncWithSmallPrefetchChunks() throws Exception { + try (AsyncResultSet resultSet = + getClient(dialect.dialect) + .singleUse() + .executeQueryAsync( + Statement.of( + "SELECT Key, Data, Fingerprint, Size FROM " + TABLE_NAME + " ORDER BY Key"), + Options.prefetchChunks(1))) { + validateAsync(resultSet); + } + } + + private void validateAsync(AsyncResultSet resultSet) throws Exception { + ExecutorService executor = Executors.newSingleThreadExecutor(); + try { + AtomicInteger rowCount = new AtomicInteger(); + ApiFuture future = + resultSet.setCallback( + executor, + ready -> { + while (true) { + switch (ready.tryNext()) { + case OK: + int index = rowCount.getAndIncrement(); + assertEquals(index, ready.getLong(0)); + ByteArray data = ready.getBytes(1); + assertEquals(ready.getLong(3), data.length()); + assertEquals(ready.getLong(2), hasher.hashBytes(data.toByteArray()).asLong()); + assertTrue(rowCount.get() <= numRows); + break; + case NOT_READY: + return CallbackResponse.CONTINUE; + case DONE: + return CallbackResponse.DONE; + } + } + }); + future.get(120, TimeUnit.SECONDS); + assertEquals(numRows, rowCount.get()); + } finally { + executor.shutdown(); + assertTrue(executor.awaitTermination(60, TimeUnit.SECONDS)); + } + } }