diff --git a/spring-messaging/src/main/java/org/springframework/messaging/handler/annotation/reactive/PayloadMethodArgumentResolver.java b/spring-messaging/src/main/java/org/springframework/messaging/handler/annotation/reactive/PayloadMethodArgumentResolver.java index 0dffe73b8b..c36a99ba97 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/handler/annotation/reactive/PayloadMethodArgumentResolver.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/handler/annotation/reactive/PayloadMethodArgumentResolver.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2019 the original author or authors. + * Copyright 2002-2021 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -38,6 +38,7 @@ import org.springframework.core.annotation.AnnotationUtils; import org.springframework.core.codec.Decoder; import org.springframework.core.codec.DecodingException; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.lang.Nullable; import org.springframework.messaging.Message; import org.springframework.messaging.MessageHeaders; @@ -232,6 +233,7 @@ public class PayloadMethodArgumentResolver implements HandlerMethodArgumentResol if (decoder.canDecode(elementType, mimeType)) { if (adapter != null && adapter.isMultiValue()) { Flux flux = content + .filter(this::nonEmptyDataBuffer) .map(buffer -> decoder.decode(buffer, elementType, mimeType, hints)) .onErrorResume(ex -> Flux.error(handleReadError(parameter, message, ex))); if (isContentRequired) { @@ -245,6 +247,7 @@ public class PayloadMethodArgumentResolver implements HandlerMethodArgumentResol else { // Single-value (with or without reactive type wrapper) Mono mono = content.next() + .filter(this::nonEmptyDataBuffer) .map(buffer -> decoder.decode(buffer, elementType, mimeType, hints)) .onErrorResume(ex -> Mono.error(handleReadError(parameter, message, ex))); if (isContentRequired) { @@ -262,6 +265,14 @@ public class PayloadMethodArgumentResolver implements HandlerMethodArgumentResol message, parameter, "Cannot decode to [" + targetType + "]" + message)); } + private boolean nonEmptyDataBuffer(DataBuffer buffer) { + if (buffer.readableByteCount() > 0) { + return true; + } + DataBufferUtils.release(buffer); + return false; + } + private Throwable handleReadError(MethodParameter parameter, Message message, Throwable ex) { return ex instanceof DecodingException ? new MethodArgumentResolutionException(message, parameter, "Failed to read HTTP message", ex) : ex; diff --git a/spring-messaging/src/test/java/org/springframework/messaging/rsocket/RSocketClientToServerIntegrationTests.java b/spring-messaging/src/test/java/org/springframework/messaging/rsocket/RSocketClientToServerIntegrationTests.java index 8dfd587ac9..12e416645d 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/rsocket/RSocketClientToServerIntegrationTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/rsocket/RSocketClientToServerIntegrationTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2020 the original author or authors. + * Copyright 2002-2021 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -19,7 +19,6 @@ package org.springframework.messaging.rsocket; import java.time.Duration; import java.util.concurrent.atomic.AtomicInteger; -import io.rsocket.Payload; import io.rsocket.RSocket; import io.rsocket.SocketAcceptor; import io.rsocket.core.RSocketServer; @@ -43,6 +42,7 @@ import org.springframework.context.annotation.Configuration; import org.springframework.messaging.handler.annotation.Header; import org.springframework.messaging.handler.annotation.MessageExceptionHandler; import org.springframework.messaging.handler.annotation.MessageMapping; +import org.springframework.messaging.handler.annotation.Payload; import org.springframework.messaging.rsocket.annotation.ConnectMapping; import org.springframework.messaging.rsocket.annotation.support.RSocketMessageHandler; import org.springframework.stereotype.Controller; @@ -164,6 +164,12 @@ public class RSocketClientToServerIntegrationTests { .verify(Duration.ofSeconds(5)); } + @Test // gh-26344 + public void echoChannelWithEmptyInput() { + Flux result = requester.route("echo-channel-empty").data(Flux.empty()).retrieveFlux(String.class); + StepVerifier.create(result).verifyComplete(); + } + @Test public void metadataPush() { Flux.just("bar", "baz") @@ -254,6 +260,11 @@ public class RSocketClientToServerIntegrationTests { return payloads.delayElements(Duration.ofMillis(10)).map(payload -> payload + " async"); } + @MessageMapping("echo-channel-empty") + Flux echoChannelEmpty(@Payload(required = false) Flux payloads) { + return payloads.map(payload -> payload + " echoed"); + } + @MessageMapping("thrown-exception") Mono handleAndThrow(String payload) { throw new IllegalArgumentException("Invalid input error"); @@ -338,29 +349,29 @@ public class RSocketClientToServerIntegrationTests { } @Override - public Mono fireAndForget(Payload payload) { + public Mono fireAndForget(io.rsocket.Payload payload) { return this.delegate.fireAndForget(payload) .doOnSuccess(aVoid -> this.fireAndForgetCount.incrementAndGet()); } @Override - public Mono metadataPush(Payload payload) { + public Mono metadataPush(io.rsocket.Payload payload) { return this.delegate.metadataPush(payload) .doOnSuccess(aVoid -> this.metadataPushCount.incrementAndGet()); } @Override - public Mono requestResponse(Payload payload) { + public Mono requestResponse(io.rsocket.Payload payload) { return this.delegate.requestResponse(payload); } @Override - public Flux requestStream(Payload payload) { + public Flux requestStream(io.rsocket.Payload payload) { return this.delegate.requestStream(payload); } @Override - public Flux requestChannel(Publisher payloads) { + public Flux requestChannel(Publisher payloads) { return this.delegate.requestChannel(payloads); } }