Reject null for non-optional arguments

Closes gh-33339
This commit is contained in:
Olga Maciaszek-Sharma
2024-08-05 15:30:10 +02:00
committed by rstoyanchev
parent 4ac4c1b868
commit 51de84e148
14 changed files with 378 additions and 81 deletions

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2022 the original author or authors.
* Copyright 2002-2024 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.
@@ -29,6 +29,7 @@ import org.springframework.util.Assert;
* annotated arguments.
*
* @author Rossen Stoyanchev
* @author Olga Maciaszek-Sharma
* @since 6.0
*/
public class PayloadArgumentResolver implements RSocketServiceArgumentResolver {
@@ -54,25 +55,28 @@ public class PayloadArgumentResolver implements RSocketServiceArgumentResolver {
return false;
}
if (argument != null) {
ReactiveAdapter reactiveAdapter = this.reactiveAdapterRegistry.getAdapter(parameter.getParameterType());
if (reactiveAdapter == null) {
requestValues.setPayloadValue(argument);
}
else {
MethodParameter nestedParameter = parameter.nested();
String message = "Async type for @Payload should produce value(s)";
Assert.isTrue(nestedParameter.getNestedParameterType() != Void.class, message);
Assert.isTrue(!reactiveAdapter.isNoValue(), message);
requestValues.setPayload(
reactiveAdapter.toPublisher(argument),
ParameterizedTypeReference.forType(nestedParameter.getNestedGenericParameterType()));
}
if (argument == null) {
boolean required = (annot == null || annot.required()) && !parameter.isOptional();
Assert.isTrue(!required, () -> "Missing payload");
return true;
}
ReactiveAdapter reactiveAdapter = this.reactiveAdapterRegistry
.getAdapter(parameter.getParameterType());
if (reactiveAdapter == null) {
requestValues.setPayloadValue(argument);
}
else {
MethodParameter nestedParameter = parameter.nested();
String message = "Async type for @Payload should produce value(s)";
Assert.isTrue(nestedParameter.getNestedParameterType() != Void.class, message);
Assert.isTrue(!reactiveAdapter.isNoValue(), message);
requestValues.setPayload(
reactiveAdapter.toPublisher(argument),
ParameterizedTypeReference.forType(nestedParameter.getNestedGenericParameterType()));
}
return true;
}
}

View File

@@ -25,6 +25,7 @@ import reactor.core.publisher.Mono;
import org.springframework.core.MethodParameter;
import org.springframework.core.ParameterizedTypeReference;
import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.lang.Nullable;
import org.springframework.messaging.handler.annotation.Payload;
import static org.assertj.core.api.Assertions.assertThat;
@@ -34,7 +35,9 @@ import static org.assertj.core.api.Assertions.assertThatIllegalArgumentException
* Tests for {@link PayloadArgumentResolver}.
*
* @author Rossen Stoyanchev
* @author Olga Maciaszek-Sharma
*/
@SuppressWarnings("DataFlowIssue")
class PayloadArgumentResolverTests extends RSocketServiceArgumentResolverTestSupport {
@Override
@@ -47,9 +50,7 @@ class PayloadArgumentResolverTests extends RSocketServiceArgumentResolverTestSup
String payload = "payloadValue";
boolean resolved = execute(payload, initMethodParameter(Service.class, "execute", 0));
assertThat(resolved).isTrue();
assertThat(getRequestValues().getPayloadValue()).isEqualTo(payload);
assertThat(getRequestValues().getPayload()).isNull();
assertPayload(resolved, payload);
}
@Test
@@ -57,10 +58,7 @@ class PayloadArgumentResolverTests extends RSocketServiceArgumentResolverTestSup
Mono<String> payloadMono = Mono.just("payloadValue");
boolean resolved = execute(payloadMono, initMethodParameter(Service.class, "executeMono", 0));
assertThat(resolved).isTrue();
assertThat(getRequestValues().getPayloadValue()).isNull();
assertThat(getRequestValues().getPayload()).isSameAs(payloadMono);
assertThat(getRequestValues().getPayloadElementType()).isEqualTo(new ParameterizedTypeReference<String>() {});
assertPayloadMono(resolved, payloadMono);
}
@Test
@@ -92,7 +90,7 @@ class PayloadArgumentResolverTests extends RSocketServiceArgumentResolverTestSup
}
@Test
void notRequestBody() {
void notPayload() {
MethodParameter parameter = initMethodParameter(Service.class, "executeNotAnnotated", 0);
boolean resolved = execute("value", parameter);
@@ -100,23 +98,69 @@ class PayloadArgumentResolverTests extends RSocketServiceArgumentResolverTestSup
}
@Test
void ignoreNull() {
boolean resolved = execute(null, initMethodParameter(Service.class, "execute", 0));
void nullPayload() {
assertThatIllegalArgumentException()
.isThrownBy(() -> execute(null, initMethodParameter(Service.class, "execute", 0)))
.withMessage("Missing payload");
assertThatIllegalArgumentException()
.isThrownBy(() -> execute(null, initMethodParameter(Service.class, "executeMono", 0)))
.withMessage("Missing payload");
}
@Test
void nullPayloadWithNullable() {
boolean resolved = execute(null, initMethodParameter(Service.class, "executeNullable", 0));
assertNullValues(resolved);
boolean resolvedMono = execute(null, initMethodParameter(Service.class, "executeNullableMono", 0));
assertNullValues(resolvedMono);
}
@Test
void nullPayloadWithNotRequired() {
boolean resolved = execute(null, initMethodParameter(Service.class, "executeNotRequired", 0));
assertNullValues(resolved);
boolean resolvedMono = execute(null, initMethodParameter(Service.class, "executeNotRequiredMono", 0));
assertNullValues(resolvedMono);
}
private void assertPayload(boolean resolved, String payload) {
assertThat(resolved).isTrue();
assertThat(getRequestValues().getPayloadValue()).isEqualTo(payload);
assertThat(getRequestValues().getPayload()).isNull();
}
private void assertPayloadMono(boolean resolved, Mono<String> payloadMono) {
assertThat(resolved).isTrue();
assertThat(getRequestValues().getPayloadValue()).isNull();
assertThat(getRequestValues().getPayload()).isSameAs(payloadMono);
assertThat(getRequestValues().getPayloadElementType()).isEqualTo(new ParameterizedTypeReference<String>() { });
}
private void assertNullValues(boolean resolved) {
assertThat(resolved).isTrue();
assertThat(getRequestValues().getPayloadValue()).isNull();
assertThat(getRequestValues().getPayload()).isNull();
assertThat(getRequestValues().getPayloadElementType()).isNull();
}
@SuppressWarnings("unused")
@SuppressWarnings({"unused"})
private interface Service {
void execute(@Payload String body);
void executeNotRequired(@Payload(required = false) String body);
void executeNullable(@Nullable @Payload String body);
void executeMono(@Payload Mono<String> body);
void executeNullableMono(@Nullable @Payload Mono<String> body);
void executeNotRequiredMono(@Payload(required = false) Mono<String> body);
void executeSingle(@Payload Single<String> body);
void executeMonoVoid(@Payload Mono<Void> body);