Skip to content

Commit a896e5b

Browse files
committed
Fix remote OOM via delayed deserialization for unary/server-streaming calls
Protobuf unknown-field or repeated field amplification can lead to remote OOM if an attacker sends a unary request but holds the stream open without half-closing. This change delays the deserialization of incoming messages for calls where the client sends at most one message (Unary and Server Streaming) until the client actually half-closes the stream (sends END_STREAM). If the call is cancelled before half-close, the buffered raw message is discarded without being deserialized, preventing the memory explosion.
1 parent e9a8c2b commit a896e5b

2 files changed

Lines changed: 148 additions & 16 deletions

File tree

core/src/main/java/io/grpc/internal/ServerCallImpl.java

Lines changed: 65 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@
3434
import io.grpc.CompressorRegistry;
3535
import io.grpc.Context;
3636
import io.grpc.DecompressorRegistry;
37+
import io.grpc.Detachable;
3738
import io.grpc.InternalDecompressorRegistry;
3839
import io.grpc.InternalStatus;
3940
import io.grpc.Metadata;
@@ -45,6 +46,9 @@
4546
import io.perfmark.PerfMark;
4647
import io.perfmark.Tag;
4748
import io.perfmark.TaskCloseable;
49+
import java.io.ByteArrayInputStream;
50+
import java.io.ByteArrayOutputStream;
51+
import java.io.IOException;
4852
import java.io.InputStream;
4953
import java.util.logging.Level;
5054
import java.util.logging.Logger;
@@ -288,6 +292,7 @@ static final class ServerStreamListenerImpl<ReqT> implements ServerStreamListene
288292
private final ServerCallImpl<ReqT, ?> call;
289293
private final ServerCall.Listener<ReqT> listener;
290294
private final Context.CancellableContext context;
295+
private InputStream delayedMessage;
291296

292297
public ServerStreamListenerImpl(
293298
ServerCallImpl<ReqT, ?> call, ServerCall.Listener<ReqT> listener,
@@ -320,6 +325,20 @@ public void messagesAvailable(MessageProducer producer) {
320325
}
321326
}
322327

328+
private static InputStream bufferMessage(InputStream is) throws IOException {
329+
if (is instanceof Detachable) {
330+
return ((Detachable) is).detach();
331+
}
332+
// Fallback: copy to byte array
333+
ByteArrayOutputStream baos = new ByteArrayOutputStream();
334+
byte[] buffer = new byte[4096];
335+
int bytesRead;
336+
while ((bytesRead = is.read(buffer)) != -1) {
337+
baos.write(buffer, 0, bytesRead);
338+
}
339+
return new ByteArrayInputStream(baos.toByteArray());
340+
}
341+
323342
@SuppressWarnings("Finally") // The code avoids suppressing the exception thrown from try
324343
private void messagesAvailableInternal(final MessageProducer producer) {
325344
if (call.cancelled) {
@@ -330,13 +349,32 @@ private void messagesAvailableInternal(final MessageProducer producer) {
330349
InputStream message;
331350
try {
332351
while ((message = producer.next()) != null) {
333-
try {
334-
listener.onMessage(call.method.parseRequest(message));
335-
} catch (Throwable t) {
336-
GrpcUtil.closeQuietly(message);
337-
throw t;
352+
if (call.method.getType().clientSendsOneMessage()) {
353+
if (delayedMessage != null) {
354+
GrpcUtil.closeQuietly(message);
355+
call.close(
356+
Status.INTERNAL.withDescription("Too many requests"),
357+
new Metadata());
358+
GrpcUtil.closeQuietly(delayedMessage);
359+
delayedMessage = null;
360+
return;
361+
}
362+
try {
363+
delayedMessage = bufferMessage(message);
364+
} catch (Throwable t) {
365+
GrpcUtil.closeQuietly(message);
366+
throw t;
367+
}
368+
message.close();
369+
} else {
370+
try {
371+
listener.onMessage(call.method.parseRequest(message));
372+
} catch (Throwable t) {
373+
GrpcUtil.closeQuietly(message);
374+
throw t;
375+
}
376+
message.close();
338377
}
339-
message.close();
340378
}
341379
} catch (Throwable t) {
342380
GrpcUtil.closeQuietly(producer);
@@ -353,6 +391,23 @@ public void halfClosed() {
353391
return;
354392
}
355393

394+
if (delayedMessage != null) {
395+
InputStream message = delayedMessage;
396+
delayedMessage = null;
397+
try {
398+
listener.onMessage(call.method.parseRequest(message));
399+
} catch (Throwable t) {
400+
GrpcUtil.closeQuietly(message);
401+
Throwables.throwIfUnchecked(t);
402+
throw new RuntimeException(t);
403+
}
404+
try {
405+
message.close();
406+
} catch (IOException e) {
407+
throw new RuntimeException(e);
408+
}
409+
}
410+
356411
listener.onHalfClose();
357412
}
358413
}
@@ -366,6 +421,10 @@ public void closed(Status status) {
366421
}
367422

368423
private void closedInternal(Status status) {
424+
if (delayedMessage != null) {
425+
GrpcUtil.closeQuietly(delayedMessage);
426+
delayedMessage = null;
427+
}
369428
Throwable cancelCause = null;
370429
try {
371430
if (status.isOk()) {

core/src/test/java/io/grpc/internal/ServerCallImplTest.java

Lines changed: 83 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,23 @@ public class ServerCallImplTest {
8484

8585
private static final MethodDescriptor<Long, Long> CLIENT_STREAMING_METHOD =
8686
MethodDescriptor.<Long, Long>newBuilder()
87-
.setType(MethodType.UNARY)
87+
.setType(MethodType.CLIENT_STREAMING)
88+
.setFullMethodName("service/method")
89+
.setRequestMarshaller(new LongMarshaller())
90+
.setResponseMarshaller(new LongMarshaller())
91+
.build();
92+
93+
private static final MethodDescriptor<Long, Long> BIDI_STREAMING_METHOD =
94+
MethodDescriptor.<Long, Long>newBuilder()
95+
.setType(MethodType.BIDI_STREAMING)
96+
.setFullMethodName("service/method")
97+
.setRequestMarshaller(new LongMarshaller())
98+
.setResponseMarshaller(new LongMarshaller())
99+
.build();
100+
101+
private static final MethodDescriptor<Long, Long> SERVER_STREAMING_METHOD =
102+
MethodDescriptor.<Long, Long>newBuilder()
103+
.setType(MethodType.SERVER_STREAMING)
88104
.setFullMethodName("service/method")
89105
.setRequestMarshaller(new LongMarshaller())
90106
.setResponseMarshaller(new LongMarshaller())
@@ -456,40 +472,97 @@ public void streamListener_onReady_onlyOnce() {
456472
}
457473

458474
@Test
459-
public void streamListener_messageRead() {
475+
public void streamListener_messageRead_unary_delayed() {
460476
ServerStreamListenerImpl<Long> streamListener =
461477
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
462478
streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L)));
463479

480+
// Message should not be delivered yet
481+
verify(callListener, never()).onMessage(any(Long.class));
482+
483+
streamListener.halfClosed();
484+
485+
// Now it should be delivered
464486
verify(callListener).onMessage(1234L);
487+
verify(callListener).onHalfClose();
465488
}
466489

467490
@Test
468-
public void streamListener_messageRead_onlyOnce() {
491+
public void streamListener_messageRead_serverStreaming_delayed() {
492+
call = new ServerCallImpl<>(stream, SERVER_STREAMING_METHOD, requestHeaders, context,
493+
DecompressorRegistry.getDefaultInstance(), CompressorRegistry.getDefaultInstance(),
494+
serverCallTracer, PerfMark.createTag());
495+
ServerStreamListenerImpl<Long> streamListener =
496+
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
497+
streamListener.messagesAvailable(new SingleMessageProducer(SERVER_STREAMING_METHOD.streamRequest(1234L)));
498+
499+
// Message should not be delivered yet
500+
verify(callListener, never()).onMessage(any(Long.class));
501+
502+
streamListener.halfClosed();
503+
504+
// Now it should be delivered
505+
verify(callListener).onMessage(1234L);
506+
verify(callListener).onHalfClose();
507+
}
508+
509+
@Test
510+
public void streamListener_messageRead_bidi_notDelayed() {
511+
call = new ServerCallImpl<>(stream, BIDI_STREAMING_METHOD, requestHeaders, context,
512+
DecompressorRegistry.getDefaultInstance(), CompressorRegistry.getDefaultInstance(),
513+
serverCallTracer, PerfMark.createTag());
514+
ServerStreamListenerImpl<Long> streamListener =
515+
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
516+
streamListener.messagesAvailable(new SingleMessageProducer(BIDI_STREAMING_METHOD.streamRequest(1234L)));
517+
518+
// Message should be delivered immediately
519+
verify(callListener).onMessage(1234L);
520+
verify(callListener, never()).onHalfClose();
521+
}
522+
523+
@Test
524+
public void streamListener_messageRead_unary_tooManyRequests() {
469525
ServerStreamListenerImpl<Long> streamListener =
470526
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
471527
streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L)));
472-
// canceling the call should short circuit future halfClosed() calls.
473-
streamListener.closed(Status.CANCELLED);
474528

529+
// Sending second message should fail
530+
streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(5678L)));
531+
532+
verify(stream).close(any(Status.class), any(Metadata.class));
533+
verify(callListener, never()).onMessage(any(Long.class));
534+
}
535+
536+
@Test
537+
public void streamListener_messageRead_onlyOnce_unary() {
538+
ServerStreamListenerImpl<Long> streamListener =
539+
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
475540
streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L)));
541+
542+
// canceling the call should clean up and prevent delivery
543+
streamListener.closed(Status.CANCELLED);
476544

477-
verify(callListener).onMessage(1234L);
545+
streamListener.halfClosed();
546+
547+
verify(callListener, never()).onMessage(any(Long.class));
478548
}
479549

480550
@Test
481-
public void streamListener_unexpectedRuntimeException() {
551+
public void streamListener_unexpectedRuntimeException_unary() {
482552
ServerStreamListenerImpl<Long> streamListener =
483553
new ServerCallImpl.ServerStreamListenerImpl<>(call, callListener, context);
484554
doThrow(new RuntimeException("unexpected exception"))
485555
.when(callListener)
486556
.onMessage(any(Long.class));
487557

488-
InputStream inputStream = UNARY_METHOD.streamRequest(1234L);
558+
streamListener.messagesAvailable(new SingleMessageProducer(UNARY_METHOD.streamRequest(1234L)));
559+
560+
// Exception should not be thrown yet because deserialization/delivery is delayed
561+
verify(callListener, never()).onMessage(any(Long.class));
489562

490-
SingleMessageProducer producer = new SingleMessageProducer(inputStream);
563+
// It should be thrown during halfClosed
491564
RuntimeException e = assertThrows(RuntimeException.class,
492-
() -> streamListener.messagesAvailable(producer));
565+
() -> streamListener.halfClosed());
493566
assertThat(e).hasMessageThat().isEqualTo("unexpected exception");
494567
}
495568

0 commit comments

Comments
 (0)