diff --git a/core/src/main/java/io/grpc/internal/MessageDeframer.java b/core/src/main/java/io/grpc/internal/MessageDeframer.java index f388c006e97..dbb0d88b87a 100644 --- a/core/src/main/java/io/grpc/internal/MessageDeframer.java +++ b/core/src/main/java/io/grpc/internal/MessageDeframer.java @@ -437,14 +437,102 @@ private InputStream getCompressedBody() { .asRuntimeException(); } - try { - // Enforce the maxMessageSize limit on the returned stream. - InputStream unlimitedStream = - decompressor.decompress(ReadableBuffers.openStream(nextFrame, true)); - return new SizeEnforcingInputStream( - unlimitedStream, maxInboundMessageSize, statsTraceCtx); - } catch (IOException e) { - throw new RuntimeException(e); + return new LazyDecompressingInputStream( + ReadableBuffers.openStream(nextFrame, true), + maxInboundMessageSize, + statsTraceCtx, + decompressor); + } + + /** + * An {@link InputStream} that delays decompressing a compressed frame until data is first read. + */ + @VisibleForTesting + static final class LazyDecompressingInputStream extends FilterInputStream { + private final Decompressor decompressor; + private final int maxMessageSize; + private final StatsTraceContext statsTraceCtx; + private boolean initialized; + private boolean closed; + + LazyDecompressingInputStream( + InputStream rawStream, + int maxMessageSize, + StatsTraceContext statsTraceCtx, + Decompressor decompressor) { + super(rawStream); + this.decompressor = decompressor; + this.maxMessageSize = maxMessageSize; + this.statsTraceCtx = statsTraceCtx; + } + + private synchronized void ensureInitialized() throws IOException { + if (closed) { + throw new IOException("Stream closed"); + } + if (!initialized) { + InputStream decompressed = decompressor.decompress(in); + in = new SizeEnforcingInputStream(decompressed, maxMessageSize, statsTraceCtx); + initialized = true; + } + } + + @Override + public int read() throws IOException { + ensureInitialized(); + return super.read(); + } + + @Override + public int read(byte[] b, int off, int len) throws IOException { + ensureInitialized(); + return super.read(b, off, len); + } + + @Override + public long skip(long n) throws IOException { + ensureInitialized(); + return super.skip(n); + } + + @Override + public int available() throws IOException { + ensureInitialized(); + return super.available(); + } + + @Override + public synchronized void close() throws IOException { + if (!closed) { + closed = true; + super.close(); + } + } + + @Override + public synchronized void mark(int readlimit) { + try { + ensureInitialized(); + super.mark(readlimit); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public synchronized void reset() throws IOException { + ensureInitialized(); + super.reset(); + } + + @Override + public boolean markSupported() { + try { + ensureInitialized(); + return super.markSupported(); + } catch (IOException e) { + throw new RuntimeException(e); + } } } diff --git a/core/src/main/java/io/grpc/internal/ServerCallImpl.java b/core/src/main/java/io/grpc/internal/ServerCallImpl.java index e224384ce8f..823f528be98 100644 --- a/core/src/main/java/io/grpc/internal/ServerCallImpl.java +++ b/core/src/main/java/io/grpc/internal/ServerCallImpl.java @@ -288,6 +288,7 @@ static final class ServerStreamListenerImpl implements ServerStreamListene private final ServerCallImpl call; private final ServerCall.Listener listener; private final Context.CancellableContext context; + private InputStream delayedMessage; public ServerStreamListenerImpl( ServerCallImpl call, ServerCall.Listener listener, @@ -330,13 +331,27 @@ private void messagesAvailableInternal(final MessageProducer producer) { InputStream message; try { while ((message = producer.next()) != null) { - try { - listener.onMessage(call.method.parseRequest(message)); - } catch (Throwable t) { - GrpcUtil.closeQuietly(message); - throw t; + // TODO: Consider forcing this check to be done in the transport (MessageDeframer) + // https://github.com/grpc/grpc-java/pull/13004/changes#r3939373996 + if (call.method.getType().clientSendsOneMessage()) { + if (delayedMessage != null) { + GrpcUtil.closeQuietly(message); + call.stream.cancel(Status.INTERNAL.withDescription("Too many requests")); + GrpcUtil.closeQuietly(delayedMessage); + delayedMessage = null; + call.cancelled = true; + return; + } + delayedMessage = message; + } else { + try { + listener.onMessage(call.method.parseRequest(message)); + } catch (Throwable t) { + GrpcUtil.closeQuietly(message); + throw t; + } + message.close(); } - message.close(); } } catch (Throwable t) { GrpcUtil.closeQuietly(producer); @@ -353,6 +368,19 @@ public void halfClosed() { return; } + if (delayedMessage != null) { + InputStream message = delayedMessage; + delayedMessage = null; + try { + listener.onMessage(call.method.parseRequest(message)); + } catch (Throwable t) { + GrpcUtil.closeQuietly(message); + Throwables.throwIfUnchecked(t); + throw new RuntimeException(t); + } + GrpcUtil.closeQuietly(message); + } + listener.onHalfClose(); } } @@ -366,6 +394,10 @@ public void closed(Status status) { } private void closedInternal(Status status) { + if (delayedMessage != null) { + GrpcUtil.closeQuietly(delayedMessage); + delayedMessage = null; + } Throwable cancelCause = null; try { if (status.isOk()) { diff --git a/core/src/test/java/io/grpc/internal/MessageDeframerTest.java b/core/src/test/java/io/grpc/internal/MessageDeframerTest.java index 54758bc096f..ea6b83bb6e6 100644 --- a/core/src/test/java/io/grpc/internal/MessageDeframerTest.java +++ b/core/src/test/java/io/grpc/internal/MessageDeframerTest.java @@ -19,6 +19,8 @@ import static com.google.common.truth.Truth.assertThat; import static io.grpc.internal.GrpcUtil.DEFAULT_MAX_MESSAGE_SIZE; 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; @@ -35,9 +37,11 @@ import com.google.common.io.ByteStreams; import com.google.common.primitives.Bytes; import io.grpc.Codec; +import io.grpc.Decompressor; import io.grpc.InternalChannelz.TransportStats; import io.grpc.StatusRuntimeException; import io.grpc.StreamTracer; +import io.grpc.internal.MessageDeframer.LazyDecompressingInputStream; import io.grpc.internal.MessageDeframer.Listener; import io.grpc.internal.MessageDeframer.SizeEnforcingInputStream; import io.grpc.internal.testing.TestStreamTracer.TestBaseStreamTracer; @@ -52,6 +56,7 @@ import java.util.List; import java.util.Locale; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.zip.GZIPOutputStream; import org.junit.Before; import org.junit.Test; @@ -313,6 +318,75 @@ public void compressed() { verifyNoMoreInteractions(listener); } + @Test + public void compressed_lazyDecompression() throws IOException { + final AtomicBoolean decompressCalled = new AtomicBoolean(false); + Decompressor countingDecompressor = new Decompressor() { + @Override + public String getMessageEncoding() { + return "gzip"; + } + + @Override + public InputStream decompress(InputStream is) throws IOException { + decompressCalled.set(true); + return new Codec.Gzip().decompress(is); + } + }; + + deframer = new MessageDeframer(listener, countingDecompressor, DEFAULT_MAX_MESSAGE_SIZE, + statsTraceCtx, transportTracer); + deframer.request(1); + + byte[] payload = compress(new byte[1000]); + byte[] header = new byte[]{1, 0, 0, 0, (byte) payload.length}; + deframer.deframe(buffer(Bytes.concat(header, payload))); + + verify(listener).messagesAvailable(producer.capture()); + InputStream stream = producer.getValue().next(); + assertNotNull(stream); + + // Decompressor should not be invoked before bytes are read + assertFalse(decompressCalled.get()); + + // Reading a byte triggers decompression + assertEquals(0, stream.read()); + assertTrue(decompressCalled.get()); + } + + @Test + public void compressed_closeWithoutReading_noDecompression() throws IOException { + final AtomicBoolean decompressCalled = new AtomicBoolean(false); + Decompressor countingDecompressor = new Decompressor() { + @Override + public String getMessageEncoding() { + return "gzip"; + } + + @Override + public InputStream decompress(InputStream is) throws IOException { + decompressCalled.set(true); + return new Codec.Gzip().decompress(is); + } + }; + + deframer = new MessageDeframer(listener, countingDecompressor, DEFAULT_MAX_MESSAGE_SIZE, + statsTraceCtx, transportTracer); + deframer.request(1); + + byte[] payload = compress(new byte[1000]); + byte[] header = new byte[]{1, 0, 0, 0, (byte) payload.length}; + deframer.deframe(buffer(Bytes.concat(header, payload))); + + verify(listener).messagesAvailable(producer.capture()); + InputStream stream = producer.getValue().next(); + assertNotNull(stream); + + // Closing without reading should not decompress + stream.close(); + assertFalse(decompressCalled.get()); + } + @Test public void deliverIsReentrantSafe() { doAnswer( @@ -493,6 +567,79 @@ public void sizeEnforcingInputStream_markReset() throws IOException { } } + @RunWith(JUnit4.class) + public static class LazyDecompressingInputStreamTests { + private TestBaseStreamTracer tracer = new TestBaseStreamTracer(); + private StatsTraceContext statsTraceCtx = new StatsTraceContext(new StreamTracer[]{tracer}); + + @Test + public void lazyDecompressingInputStream_doesNotInitializeUntilRead() throws IOException { + final AtomicBoolean decompressCalled = new AtomicBoolean(false); + Decompressor countingDecompressor = new Decompressor() { + @Override + public String getMessageEncoding() { + return "gzip"; + } + + @Override + public InputStream decompress(InputStream is) throws IOException { + decompressCalled.set(true); + return new Codec.Gzip().decompress(is); + } + }; + + ByteArrayInputStream in = + new ByteArrayInputStream(compress("hello".getBytes(StandardCharsets.UTF_8))); + LazyDecompressingInputStream stream = new LazyDecompressingInputStream( + in, 100, statsTraceCtx, countingDecompressor); + + assertFalse(decompressCalled.get()); + byte[] buf = new byte[5]; + int read = stream.read(buf); + assertEquals(5, read); + assertEquals("hello", new String(buf, StandardCharsets.UTF_8)); + assertTrue(decompressCalled.get()); + stream.close(); + } + + @Test + public void lazyDecompressingInputStream_closeWithoutRead() throws IOException { + final AtomicBoolean decompressCalled = new AtomicBoolean(false); + final AtomicBoolean inClosed = new AtomicBoolean(false); + Decompressor countingDecompressor = new Decompressor() { + @Override + public String getMessageEncoding() { + return "gzip"; + } + + @Override + public InputStream decompress(InputStream is) throws IOException { + decompressCalled.set(true); + return new Codec.Gzip().decompress(is); + } + }; + + ByteArrayInputStream in = + new ByteArrayInputStream(compress("hello".getBytes(StandardCharsets.UTF_8))) { + @Override + public void close() throws IOException { + inClosed.set(true); + super.close(); + } + }; + LazyDecompressingInputStream stream = new LazyDecompressingInputStream( + in, 100, statsTraceCtx, countingDecompressor); + + assertFalse(decompressCalled.get()); + stream.close(); + assertTrue(inClosed.get()); + assertFalse(decompressCalled.get()); + + // Reading after close should throw IOException + assertThrows(IOException.class, () -> stream.read()); + } + } + /** * Verify stats were published through the tracer. * diff --git a/core/src/test/java/io/grpc/internal/ServerCallImplTest.java b/core/src/test/java/io/grpc/internal/ServerCallImplTest.java index 7394c83eab2..9eb3e5b959b 100644 --- a/core/src/test/java/io/grpc/internal/ServerCallImplTest.java +++ b/core/src/test/java/io/grpc/internal/ServerCallImplTest.java @@ -51,6 +51,7 @@ import io.grpc.internal.ServerCallImpl.ServerStreamListenerImpl; import io.perfmark.PerfMark; import java.io.ByteArrayInputStream; +import java.io.IOException; import java.io.InputStream; import java.io.InputStreamReader; import org.junit.Before; @@ -84,7 +85,23 @@ public class ServerCallImplTest { private static final MethodDescriptor CLIENT_STREAMING_METHOD = MethodDescriptor.newBuilder() - .setType(MethodType.UNARY) + .setType(MethodType.CLIENT_STREAMING) + .setFullMethodName("service/method") + .setRequestMarshaller(new LongMarshaller()) + .setResponseMarshaller(new LongMarshaller()) + .build(); + + private static final MethodDescriptor BIDI_STREAMING_METHOD = + MethodDescriptor.newBuilder() + .setType(MethodType.BIDI_STREAMING) + .setFullMethodName("service/method") + .setRequestMarshaller(new LongMarshaller()) + .setResponseMarshaller(new LongMarshaller()) + .build(); + + private static final MethodDescriptor SERVER_STREAMING_METHOD = + MethodDescriptor.newBuilder() + .setType(MethodType.SERVER_STREAMING) .setFullMethodName("service/method") .setRequestMarshaller(new LongMarshaller()) .setResponseMarshaller(new LongMarshaller()) @@ -456,41 +473,153 @@ public void streamListener_onReady_onlyOnce() { } @Test - public void streamListener_messageRead() { + public void streamListener_messageRead_unary_delayed() { ServerStreamListenerImpl streamListener = new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); - streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L))); + FakeCloseTrackerInputStream messageStream = + new FakeCloseTrackerInputStream(UNARY_METHOD.streamRequest(1234L), false); + streamListener.messagesAvailable(new SingleMessageProducer(messageStream)); + + // Message should not be delivered or closed yet + verify(callListener, never()).onMessage(any(Long.class)); + assertFalse(messageStream.closed); + + streamListener.halfClosed(); + // Now it should be delivered and closed verify(callListener).onMessage(1234L); + verify(callListener).onHalfClose(); + assertTrue(messageStream.closed); } @Test - public void streamListener_messageRead_onlyOnce() { + public void streamListener_messageRead_serverStreaming_delayed() { + call = new ServerCallImpl<>(stream, SERVER_STREAMING_METHOD, requestHeaders, context, + DecompressorRegistry.getDefaultInstance(), CompressorRegistry.getDefaultInstance(), + serverCallTracer, PerfMark.createTag()); + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + FakeCloseTrackerInputStream messageStream = + new FakeCloseTrackerInputStream( + SERVER_STREAMING_METHOD.streamRequest(1234L), false); + streamListener.messagesAvailable(new SingleMessageProducer(messageStream)); + + // Message should not be delivered or closed yet + verify(callListener, never()).onMessage(any(Long.class)); + assertFalse(messageStream.closed); + + streamListener.halfClosed(); + + // Now it should be delivered and closed + verify(callListener).onMessage(1234L); + verify(callListener).onHalfClose(); + assertTrue(messageStream.closed); + } + + @Test + public void streamListener_messageRead_bidi_notDelayed() { + call = new ServerCallImpl<>(stream, BIDI_STREAMING_METHOD, requestHeaders, context, + DecompressorRegistry.getDefaultInstance(), CompressorRegistry.getDefaultInstance(), + serverCallTracer, PerfMark.createTag()); + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + streamListener.messagesAvailable( + new SingleMessageProducer(BIDI_STREAMING_METHOD.streamRequest(1234L))); + + // Message should be delivered immediately + verify(callListener).onMessage(1234L); + verify(callListener, never()).onHalfClose(); + } + + @Test + public void streamListener_messageRead_unary_tooManyRequests() { ServerStreamListenerImpl streamListener = new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L))); - // canceling the call should short circuit future halfClosed() calls. + + // Sending second message should fail + streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(5678L))); + + verify(stream).cancel(any(Status.class)); + verify(callListener, never()).onMessage(any(Long.class)); + assertTrue(call.isCancelled()); + + // Subsequent halfClosed should be ignored because call is cancelled + streamListener.halfClosed(); + verify(callListener, never()).onHalfClose(); + + // When transport closes, onCancel is called once streamListener.closed(Status.CANCELLED); + verify(callListener).onCancel(); + assertTrue(context.isCancelled()); + } - streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L))); + @Test + public void streamListener_messageRead_onlyOnce_unary() { + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + FakeCloseTrackerInputStream messageStream = + new FakeCloseTrackerInputStream(UNARY_METHOD.streamRequest(1234L), false); + streamListener.messagesAvailable(new SingleMessageProducer(messageStream)); - verify(callListener).onMessage(1234L); + assertFalse(messageStream.closed); + + // canceling the call should clean up and prevent delivery + streamListener.closed(Status.CANCELLED); + + streamListener.halfClosed(); + + verify(callListener, never()).onMessage(any(Long.class)); + assertTrue(messageStream.closed); } @Test - public void streamListener_unexpectedRuntimeException() { + public void streamListener_unexpectedRuntimeException_unary() { ServerStreamListenerImpl streamListener = new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); doThrow(new RuntimeException("unexpected exception")) .when(callListener) .onMessage(any(Long.class)); - InputStream inputStream = UNARY_METHOD.streamRequest(1234L); + FakeCloseTrackerInputStream messageStream = + new FakeCloseTrackerInputStream(UNARY_METHOD.streamRequest(1234L), false); + + streamListener.messagesAvailable(new SingleMessageProducer(messageStream)); - SingleMessageProducer producer = new SingleMessageProducer(inputStream); + // Exception should not be thrown yet because deserialization/delivery is delayed + verify(callListener, never()).onMessage(any(Long.class)); + assertFalse(messageStream.closed); + + // It should be thrown during halfClosed RuntimeException e = assertThrows(RuntimeException.class, - () -> streamListener.messagesAvailable(producer)); + () -> streamListener.halfClosed()); assertThat(e).hasMessageThat().isEqualTo("unexpected exception"); + + // The stream should have been closed in the catch block + assertTrue(messageStream.closed); + } + + @Test + public void streamListener_halfClosed_closeException() { + ServerStreamListenerImpl streamListener = + new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context); + + FakeCloseTrackerInputStream messageStream = + new FakeCloseTrackerInputStream(UNARY_METHOD.streamRequest(1234L), true); + + streamListener.messagesAvailable(new SingleMessageProducer(messageStream)); + + // Message should not be delivered yet + verify(callListener, never()).onMessage(any(Long.class)); + assertFalse(messageStream.closed); + + // halfClosed should not throw because we use closeQuietly + streamListener.halfClosed(); + + // The message was delivered and halfClosed completed + verify(callListener).onMessage(1234L); + verify(callListener).onHalfClose(); + assertTrue(messageStream.closed); } private static class LongMarshaller implements Marshaller { @@ -508,4 +637,29 @@ public Long parse(InputStream stream) { } } } + + private static class FakeCloseTrackerInputStream extends InputStream { + boolean closed = false; + private final InputStream delegate; + private final boolean throwOnClose; + + FakeCloseTrackerInputStream(InputStream delegate, boolean throwOnClose) { + this.delegate = delegate; + this.throwOnClose = throwOnClose; + } + + @Override + public int read() throws IOException { + return delegate.read(); + } + + @Override + public void close() throws IOException { + closed = true; + if (throwOnClose) { + throw new IOException("close failed"); + } + delegate.close(); + } + } }