diff --git a/docs/src/main/asciidoc/sqs.adoc b/docs/src/main/asciidoc/sqs.adoc index d9d91c011..2b6e2fb8a 100644 --- a/docs/src/main/asciidoc/sqs.adoc +++ b/docs/src/main/asciidoc/sqs.adoc @@ -1782,9 +1782,13 @@ This allows payloads to be deserialized early in the message processing flow wit This enables accessing the deserialized payload in components such as `MessageInterceptor`, `ErrorHandler`, and `AcknowledgementResultCallback` without type headers. -The inference supports simple types, generic types like `List`, `Message`, and `List>`. +The inference supports simple types, generic types like `List`, `Message`, +`List>`, `Wrapper`, and `Message>`. Parameters annotated with `@Payload` are explicitly recognized as the payload parameter. +Generic type information is provided to payload converters through the `SmartMessageConverter` conversion hint. +Custom payload converters that only implement `MessageConverter` continue to receive the inferred raw payload class. + For polymorphic types (interfaces, `Object`, or `@SqsHandler` methods), a custom mapper is required. See <>. diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/AbstractEndpoint.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/AbstractEndpoint.java index 98cfe20c7..fa2c18507 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/AbstractEndpoint.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/AbstractEndpoint.java @@ -211,10 +211,10 @@ public void setupContainer(MessageListenerContainer container) { if (this.methodPayloadTypeInferrer != null && container instanceof AbstractMessageListenerContainer amlc) { - Class inferredType = this.methodPayloadTypeInferrer.inferPayloadType(this.method, + MethodPayloadMetadata payloadMetadata = this.methodPayloadTypeInferrer.inferPayloadMetadata(this.method, this.argumentResolvers); - if (inferredType != null) { - amlc.setPayloadDeserializationType(inferredType); + if (payloadMetadata != null) { + amlc.setPayloadDeserializationType(payloadMetadata.payloadClass(), payloadMetadata.conversionHint()); } disableDefaultPayloadTypeMapper(amlc); } diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrer.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrer.java index 8b9c33f08..3aa1d8795 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrer.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrer.java @@ -49,7 +49,15 @@ public class DefaultMethodPayloadTypeInferrer implements MethodPayloadTypeInferr @Override @Nullable - public Class inferPayloadType(Method method, List argumentResolvers) { + public Class inferPayloadType(Method method, @Nullable List argumentResolvers) { + MethodPayloadMetadata metadata = inferPayloadMetadata(method, argumentResolvers); + return metadata != null ? metadata.payloadClass() : null; + } + + @Override + @Nullable + public MethodPayloadMetadata inferPayloadMetadata(Method method, + @Nullable List argumentResolvers) { if (argumentResolvers == null || argumentResolvers.isEmpty()) { return null; } @@ -61,14 +69,14 @@ public Class inferPayloadType(Method method, List resolver.supportsParameter(parameter)); if (!supportedByNonPayloadResolver) { - return extractClass(parameter.getGenericParameterType()); + return extractMetadata(parameter); } } @@ -76,13 +84,14 @@ public Class inferPayloadType(Method method, List} by extracting the element type. - * @param type the inferred payload type - * @return the class to be used for payload conversion, or null if cannot be determined + * Extract the target class and conversion hint from the inferred payload parameter. Collection parameters represent + * batch listeners, so the conversion hint is nested to point at the payload of each individual message. Other + * parameters retain the method parameter so smart converters can recover their generic type. + * @param parameter the inferred payload method parameter + * @return the metadata to be used for payload conversion */ - @Nullable - private Class extractClass(Type type) { + private MethodPayloadMetadata extractMetadata(MethodParameter parameter) { + Type type = parameter.getGenericParameterType(); ResolvableType resolvableType = ResolvableType.forType(type); Class rawClass = resolvableType.toClass(); @@ -91,18 +100,17 @@ private Class extractClass(Type type) { Class elementClass = resolvableType.getNested(2).toClass(); // If it's a Collection of Messages (e.g., List>), go one level deeper if (Message.class.isAssignableFrom(elementClass)) { - return resolvableType.getNested(3).toClass(); + return new MethodPayloadMetadata(resolvableType.getNested(3).toClass(), parameter.nested().nested()); } - return elementClass; + return new MethodPayloadMetadata(elementClass, parameter.nested()); } // If it's a Message, unwrap to get T if (Message.class.isAssignableFrom(rawClass)) { - return resolvableType.getNested(2).toClass(); + return new MethodPayloadMetadata(resolvableType.getNested(2).toClass(), parameter); } - // For simple types, return as-is - return rawClass; + return new MethodPayloadMetadata(rawClass, parameter); } private boolean isPayloadResolver(HandlerMethodArgumentResolver resolver) { diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadMetadata.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadMetadata.java new file mode 100644 index 000000000..1e586f9a7 --- /dev/null +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadMetadata.java @@ -0,0 +1,29 @@ +/* + * Copyright 2013-2026 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. + * You may obtain a copy of the License at + * + * https://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 io.awspring.cloud.sqs.config; + +import org.jspecify.annotations.Nullable; + +/** + * Metadata inferred from a listener method payload parameter. + * @param payloadClass the raw class used as the message conversion target + * @param conversionHint an optional hint passed to a + * {@link org.springframework.messaging.converter.SmartMessageConverter} + * @author Bruno Augusto Garcia + * @since 4.1.0 + */ +public record MethodPayloadMetadata(Class payloadClass, @Nullable Object conversionHint) { +} diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadTypeInferrer.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadTypeInferrer.java index 703da460b..cc11f77a0 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadTypeInferrer.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/config/MethodPayloadTypeInferrer.java @@ -33,11 +33,29 @@ public interface MethodPayloadTypeInferrer { /** * Infer the payload class from the given method and its argument resolvers. + * @deprecated in favor of {@link #inferPayloadMetadata(Method, List)} * @param method the listener method * @param argumentResolvers the argument resolvers available for this method, may be null or empty * @return the inferred payload class, or null if it cannot be determined */ @Nullable + @Deprecated Class inferPayloadType(Method method, @Nullable List argumentResolvers); + /** + * Infer payload metadata from the given method and its argument resolvers. + *

+ * The default implementation adapts existing {@link MethodPayloadTypeInferrer} implementations by returning the + * inferred class without a conversion hint. + * @param method the listener method + * @param argumentResolvers the argument resolvers available for this method, may be null or empty + * @return the inferred payload metadata, or null if it cannot be determined + */ + @Nullable + default MethodPayloadMetadata inferPayloadMetadata(Method method, + @Nullable List argumentResolvers) { + Class payloadClass = inferPayloadType(method, argumentResolvers); + return payloadClass != null ? new MethodPayloadMetadata(payloadClass, null) : null; + } + } diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainer.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainer.java index fa5a4fee5..0d9af3ef2 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainer.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainer.java @@ -74,6 +74,9 @@ public abstract class AbstractMessageListenerContainer payloadDeserializationType; + @Nullable + private Object payloadConversionHint; + /** * Create an instance with the provided {@link ContainerOptions} * @param containerOptions the options instance. @@ -189,7 +192,21 @@ public void setPhase(int phase) { * @see io.awspring.cloud.sqs.support.converter.AbstractMessagingMessageConverter */ public void setPayloadDeserializationType(@Nullable Class payloadDeserializationType) { + setPayloadDeserializationType(payloadDeserializationType, null); + } + + /** + * Set the target type and conversion hint for payload deserialization. + *

+ * Since 4.0.0, the target type is typically inferred automatically from the {@code @SqsListener} method signature. + * A conversion hint can additionally preserve generic type information from that method. + * @param payloadDeserializationType the target type for deserialization + * @param conversionHint an optional hint for a smart message converter + */ + public void setPayloadDeserializationType(@Nullable Class payloadDeserializationType, + @Nullable Object conversionHint) { this.payloadDeserializationType = payloadDeserializationType; + this.payloadConversionHint = payloadDeserializationType != null ? conversionHint : null; } /** @@ -256,6 +273,15 @@ public Class getPayloadDeserializationType() { return this.payloadDeserializationType; } + /** + * Return the payload conversion hint, or null if not set. + * @return the payload conversion hint. + */ + @Nullable + public Object getPayloadConversionHint() { + return this.payloadConversionHint; + } + @Override public String getId() { return this.id; diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractPipelineMessageListenerContainer.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractPipelineMessageListenerContainer.java index 2e4cb675c..589b8ac6d 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractPipelineMessageListenerContainer.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/AbstractPipelineMessageListenerContainer.java @@ -172,7 +172,7 @@ protected void configureMessageSources(ContainerComponentFactory component teac -> teac.setTaskExecutor(taskExecutor)) .acceptManyIfNotNullAndInstance(getPayloadDeserializationType(), this.messageSources, AbstractMessageConvertingMessageSource.class, - (type, source) -> source.setPayloadDeserializationType(type)); + (type, source) -> source.setPayloadDeserializationType(type, getPayloadConversionHint())); doConfigureMessageSources(this.messageSources); } diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSource.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSource.java index 9f33fb6a7..1f01722ec 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSource.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSource.java @@ -85,8 +85,18 @@ protected void setupAcknowledgementForConversion(AcknowledgementCallback call * @param payloadDeserializationType the target class */ public void setPayloadDeserializationType(@Nullable Class payloadDeserializationType) { + this.setPayloadDeserializationType(payloadDeserializationType, null); + } + + /** + * Set the payload deserialization type and conversion hint. + * @param payloadDeserializationType the target class + * @param conversionHint an optional hint for a smart message converter + */ + public void setPayloadDeserializationType(@Nullable Class payloadDeserializationType, + @Nullable Object conversionHint) { ConfigUtils.INSTANCE.acceptBothIfNoneNull(payloadDeserializationType, this.messageConversionContext, - this::doConfigurePayloadTypeOnContext); + (payloadType, context) -> doConfigurePayloadTypeOnContext(payloadType, conversionHint, context)); } /** @@ -98,6 +108,21 @@ public void setPayloadDeserializationType(@Nullable Class payloadDeserializat protected void doConfigurePayloadTypeOnContext(Class payloadType, MessageConversionContext context) { } + /** + * Hook method for subclasses to configure the payload type and conversion hint on their specific + * {@link MessageConversionContext} implementation. + *

+ * The default implementation delegates to {@link #doConfigurePayloadTypeOnContext(Class, MessageConversionContext)} + * for backwards compatibility with existing subclasses. + * @param payloadType the payload type to configure + * @param conversionHint an optional hint for a smart message converter + * @param context the message conversion context + */ + protected void doConfigurePayloadTypeOnContext(Class payloadType, @Nullable Object conversionHint, + MessageConversionContext context) { + doConfigurePayloadTypeOnContext(payloadType, context); + } + @Nullable private MessageConversionContext maybeCreateConversionContext() { return this.messagingMessageConverter instanceof ContextAwareMessagingMessageConverter diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSource.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSource.java index 7541262f1..6be0ffcb6 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSource.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSource.java @@ -56,9 +56,9 @@ *

* *

- * Note that currently the payload is not converted here and is returned as String. The actual conversion to the - * {@link io.awspring.cloud.sqs.annotation.SqsListener} argument type happens on - * {@link org.springframework.messaging.handler.invocation.InvocableHandlerMethod} invocation. + * Payload conversion happens in the message source before the resulting message is emitted to the message sink and + * processing pipeline. When available, the inferred listener payload class and conversion hint are configured on the + * {@link SqsMessageConversionContext} and used by the messaging message converter during this step. *

* * @param the {@link Message} payload type. @@ -141,6 +141,14 @@ protected void doConfigurePayloadTypeOnContext(Class payloadType, MessageConv ctx -> ctx.setPayloadClass(payloadType)); } + @Override + protected void doConfigurePayloadTypeOnContext(Class payloadType, @Nullable Object conversionHint, + MessageConversionContext context) { + doConfigurePayloadTypeOnContext(payloadType, context); + ConfigUtils.INSTANCE.acceptIfInstance(context, SqsMessageConversionContext.class, + ctx -> ctx.setConversionHint(conversionHint)); + } + // @formatter:off private QueueAttributes resolveQueueAttributes() { return QueueAttributesResolver.builder() diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/AbstractMessagingMessageConverter.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/AbstractMessagingMessageConverter.java index 67f31cc2e..add55ada0 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/AbstractMessagingMessageConverter.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/AbstractMessagingMessageConverter.java @@ -26,6 +26,7 @@ import org.springframework.messaging.Message; import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.converter.MessageConverter; +import org.springframework.messaging.converter.SmartMessageConverter; import org.springframework.messaging.converter.StringMessageConverter; import org.springframework.messaging.support.MessageBuilder; import org.springframework.util.Assert; @@ -171,19 +172,23 @@ private MessageHeaders getContextHeaders(S message, MessageConversionContext con private Object convertPayload(S message, MessageHeaders messageHeaders, @Nullable MessageConversionContext context) { Message messagingMessage = MessageBuilder.createMessage(getPayloadToDeserialize(message), messageHeaders); - Class targetType = getTargetType(messagingMessage, context); - return targetType != null - ? Objects.requireNonNull(this.payloadMessageConverter.fromMessage(messagingMessage, targetType), - "payloadMessageConverter returned null payload") + Class mappedTargetType = this.payloadTypeMapper.apply(messagingMessage); + if (mappedTargetType != null) { + return convertPayload(messagingMessage, mappedTargetType, null); + } + + Class inferredTargetType = context != null ? context.getPayloadClass() : null; + return inferredTargetType != null + ? convertPayload(messagingMessage, inferredTargetType, context.getConversionHint()) : messagingMessage.getPayload(); } - @Nullable - private Class getTargetType(Message messagingMessage, @Nullable MessageConversionContext context) { - Class classFromTypeMapper = this.payloadTypeMapper.apply(messagingMessage); - return classFromTypeMapper == null && context != null && context.getPayloadClass() != null - ? context.getPayloadClass() - : classFromTypeMapper; + private Object convertPayload(Message message, Class targetType, @Nullable Object conversionHint) { + Object convertedPayload = conversionHint != null + && this.payloadMessageConverter instanceof SmartMessageConverter smartMessageConverter + ? smartMessageConverter.fromMessage(message, targetType, conversionHint) + : this.payloadMessageConverter.fromMessage(message, targetType); + return Objects.requireNonNull(convertedPayload, "payloadMessageConverter returned null payload"); } protected abstract Object getPayloadToDeserialize(S message); diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/MessageConversionContext.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/MessageConversionContext.java index 6472db098..3c4cd0ed8 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/MessageConversionContext.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/MessageConversionContext.java @@ -33,4 +33,13 @@ public interface MessageConversionContext { @Nullable Class getPayloadClass(); + /** + * An optional hint to be used by the payload conversion process. + * @return the conversion hint. + */ + @Nullable + default Object getConversionHint() { + return null; + } + } diff --git a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/SqsMessageConversionContext.java b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/SqsMessageConversionContext.java index 84ccf0e92..b20de181e 100644 --- a/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/SqsMessageConversionContext.java +++ b/spring-cloud-aws-sqs/src/main/java/io/awspring/cloud/sqs/support/converter/SqsMessageConversionContext.java @@ -48,6 +48,9 @@ public class SqsMessageConversionContext @Nullable private Class payloadClass; + @Nullable + private Object conversionHint; + @Override public void setQueueAttributes(QueueAttributes queueAttributes) { this.queueAttributes = queueAttributes; @@ -67,6 +70,10 @@ public void setPayloadClass(Class payloadClass) { this.payloadClass = payloadClass; } + public void setConversionHint(@Nullable Object conversionHint) { + this.conversionHint = conversionHint; + } + @Nullable public SqsAsyncClient getSqsAsyncClient() { return this.sqsAsyncClient; @@ -87,4 +94,10 @@ public AcknowledgementCallback getAcknowledgementCallback() { public Class getPayloadClass() { return this.payloadClass; } + + @Nullable + @Override + public Object getConversionHint() { + return this.conversionHint; + } } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/AbstractEndpointTest.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/AbstractEndpointTest.java index 551926d48..cce2814fe 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/AbstractEndpointTest.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/AbstractEndpointTest.java @@ -121,6 +121,21 @@ void shouldConfigureMapperWhenUsingDefaultMapper() throws Exception { verify(converter).setPayloadTypeMapper(any(Function.class)); } + @Test + void shouldConfigurePayloadDeserializationMetadataOnContainer() throws Exception { + Method method = TestListener.class.getMethod("handleMessage", String.class); + MethodParameter conversionHint = new MethodParameter(method, 0); + MethodPayloadMetadata metadata = new MethodPayloadMetadata(String.class, conversionHint); + endpoint.setMethod(method); + endpoint.setMethodPayloadTypeInferrer(inferrer); + + when(inferrer.inferPayloadMetadata(any(Method.class), any())).thenReturn(metadata); + + endpoint.setupContainer(container); + + verify(container).setPayloadDeserializationType(String.class, conversionHint); + } + @Test @SuppressWarnings("unchecked") void shouldNotOverrideCustomMapper() throws Exception { diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrerTest.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrerTest.java index b7df41a6c..418e1411e 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrerTest.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/config/DefaultMethodPayloadTypeInferrerTest.java @@ -257,6 +257,86 @@ void shouldHandleComplexNestedGenerics() throws Exception { assertThat(result).isEqualTo(CustomEvent.class); } + @Test + void shouldPreserveMethodParameterAsConversionHintForGenericWrapper() throws Exception { + Method method = TestMethods.class.getMethod("genericWrapperOfCustomEvent", GenericWrapper.class); + + MethodPayloadMetadata metadata = inferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isNotNull(); + assertThat(metadata.payloadClass()).isEqualTo(GenericWrapper.class); + assertThat(metadata.conversionHint()).isInstanceOfSatisfying(MethodParameter.class, methodParameter -> { + assertThat(methodParameter.getMethod()).isEqualTo(method); + assertThat(methodParameter.getParameterIndex()).isZero(); + assertThat(methodParameter.getGenericParameterType()).isEqualTo(method.getGenericParameterTypes()[0]); + }); + } + + @Test + void shouldPreserveMethodParameterAsConversionHintForMessageWithGenericCollectionPayload() throws Exception { + Method method = TestMethods.class.getMethod("messageOfListOfCustomEvents", Message.class); + + MethodPayloadMetadata metadata = inferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isNotNull(); + assertThat(metadata.payloadClass()).isEqualTo(List.class); + assertThat(metadata.conversionHint()).isInstanceOfSatisfying(MethodParameter.class, methodParameter -> { + assertThat(methodParameter.getMethod()).isEqualTo(method); + assertThat(methodParameter.getParameterIndex()).isZero(); + assertThat(methodParameter.getGenericParameterType()).isEqualTo(method.getGenericParameterTypes()[0]); + }); + } + + @Test + void shouldNestConversionHintForBatchGenericWrapper() throws Exception { + Method method = TestMethods.class.getMethod("listOfGenericWrapperOfCustomEvent", List.class); + + MethodPayloadMetadata metadata = inferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isNotNull(); + assertThat(metadata.payloadClass()).isEqualTo(GenericWrapper.class); + assertThat(metadata.conversionHint()).isInstanceOfSatisfying(MethodParameter.class, methodParameter -> { + assertThat(methodParameter.getNestingLevel()).isEqualTo(2); + assertThat(methodParameter.getNestedGenericParameterType().getTypeName()) + .isEqualTo(GenericWrapper.class.getName() + "<" + CustomEvent.class.getName() + ">"); + }); + } + + @Test + void shouldNestConversionHintPastMessageForBatchGenericWrapper() throws Exception { + Method method = TestMethods.class.getMethod("listOfMessageOfGenericWrapperOfCustomEvent", List.class); + + MethodPayloadMetadata metadata = inferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isNotNull(); + assertThat(metadata.payloadClass()).isEqualTo(GenericWrapper.class); + assertThat(metadata.conversionHint()).isInstanceOfSatisfying(MethodParameter.class, methodParameter -> { + assertThat(methodParameter.getNestingLevel()).isEqualTo(3); + assertThat(methodParameter.getNestedGenericParameterType().getTypeName()) + .isEqualTo(GenericWrapper.class.getName() + "<" + CustomEvent.class.getName() + ">"); + }); + } + + @Test + void shouldAdaptCustomInferrerToPayloadMetadataWithoutConversionHint() throws Exception { + Method method = TestMethods.class.getMethod("customEventPayload", CustomEvent.class); + MethodPayloadTypeInferrer customInferrer = (listenerMethod, argumentResolvers) -> CustomEvent.class; + + MethodPayloadMetadata metadata = customInferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isEqualTo(new MethodPayloadMetadata(CustomEvent.class, null)); + } + + @Test + void shouldReturnNullMetadataWhenCustomInferrerCannotInferPayloadType() throws Exception { + Method method = TestMethods.class.getMethod("customEventPayload", CustomEvent.class); + MethodPayloadTypeInferrer customInferrer = (listenerMethod, argumentResolvers) -> null; + + MethodPayloadMetadata metadata = customInferrer.inferPayloadMetadata(method, createNonSupportingResolvers()); + + assertThat(metadata).isNull(); + } + // ========== Various Payload Type Tests ========== @Test @@ -438,6 +518,20 @@ private List createOnlyPayloadResolvers() { return createStandardResolvers(); } + private List createNonSupportingResolvers() { + return List.of(new HandlerMethodArgumentResolver() { + @Override + public boolean supportsParameter(MethodParameter parameter) { + return false; + } + + @Override + public Object resolveArgument(MethodParameter parameter, Message message) { + return null; + } + }); + } + private List createMixedResolvers() { List resolvers = createStandardResolvers(); @@ -548,6 +642,18 @@ public void listOfMessageOfCustomEvent(List> messages) { public void collectionOfMessageOfCustomEvent(Collection> messages) { } + public void genericWrapperOfCustomEvent(GenericWrapper event) { + } + + public void messageOfListOfCustomEvents(Message> message) { + } + + public void listOfGenericWrapperOfCustomEvent(List> events) { + } + + public void listOfMessageOfGenericWrapperOfCustomEvent(List>> messages) { + } + // Message parameter public void messageParameter(Message message) { } @@ -558,6 +664,20 @@ public void onlyPayload(CustomEvent payload) { } + static class GenericWrapper { + + private T value; + + public T getValue() { + return this.value; + } + + public void setValue(T value) { + this.value = value; + } + + } + static class CustomEvent { private String id; diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsPayloadTypeInferenceIntegrationTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsPayloadTypeInferenceIntegrationTests.java index 7f44488c1..051745b01 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsPayloadTypeInferenceIntegrationTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/integration/SqsPayloadTypeInferenceIntegrationTests.java @@ -92,6 +92,14 @@ class SqsPayloadTypeInferenceIntegrationTests extends BaseSqsIntegrationTest { static final String ERROR_HANDLER_TEST_QUEUE = "error_handler_type_inference_queue"; + static final String INFERS_MESSAGE_LIST_PAYLOAD_QUEUE = "infers_message_list_payload_queue"; + + static final String INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE = "infers_generic_outer_payload_queue"; + + static final String INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE = "infers_batch_generic_outer_payload_queue"; + + static final String INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE = "infers_batch_message_generic_outer_payload_queue"; + static final String MANUAL_ACK_FACTORY = "manualAckFactory"; static final String CUSTOM_CONVERTER_FACTORY = "customConverterFactory"; @@ -105,7 +113,10 @@ static void beforeTests() { createQueue(client, ASYNC_LISTENER_QUEUE), createQueue(client, BATCH_MESSAGE_WRAPPER_QUEUE), createQueue(client, IGNORES_TYPE_HEADER_QUEUE), createQueue(client, EXPLICIT_PAYLOAD_ANNOTATION_QUEUE), createQueue(client, STRING_PAYLOAD_QUEUE), createQueue(client, CUSTOM_CONVERTER_QUEUE), - createQueue(client, ERROR_HANDLER_TEST_QUEUE)).join(); + createQueue(client, ERROR_HANDLER_TEST_QUEUE), createQueue(client, INFERS_MESSAGE_LIST_PAYLOAD_QUEUE), + createQueue(client, INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE), + createQueue(client, INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE), + createQueue(client, INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE)).join(); } @Autowired @@ -327,6 +338,93 @@ void errorHandlerShouldReceiveDeserializedPojo() throws Exception { errorHandlerPayloadTypeCollector.assertPayloadsForQueueContains(ERROR_HANDLER_TEST_QUEUE, event); } + @Test + void shouldInferListPayloadTypeFromMessageWrapper() throws Exception { + CountDownLatch ackLatch = new CountDownLatch(1); + ackCallbackPayloadTypeCollector.registerLatch(INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, ackLatch); + + List testEvents = List.of(new TestEvent("test-message-list-pojo-id-1", "test-payload-1"), + new TestEvent("test-message-list-pojo-id-2", "test-payload-2")); + Message> message = MessageBuilder.withPayload(testEvents).build(); + sqsTemplate.send(INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, message); + logger.debug("Sent message with List payload to queue {}: {}", INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, + message); + + assertThat(latchContainer.infersMessageListPayloadLatch.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(pojoCollector.receivedMessageListPayloads).hasSize(1); + List receivedPayload = pojoCollector.receivedMessageListPayloads.get(0).getPayload(); + assertThat(receivedPayload).allSatisfy(element -> assertThat(element).isInstanceOf(TestEvent.class)); + assertThat(receivedPayload).isEqualTo(testEvents); + interceptorPayloadTypeCollector.assertPayloadsForQueueContains(INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, testEvents); + assertThat(ackLatch.await(10, TimeUnit.SECONDS)).isTrue(); + ackCallbackPayloadTypeCollector.assertPayloadsForQueueContains(INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, testEvents); + } + + @Test + void shouldInferTestEventTypeFromGenericWrapper() throws Exception { + CountDownLatch ackLatch = new CountDownLatch(1); + ackCallbackPayloadTypeCollector.registerLatch(INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE, ackLatch); + + GenericWrapperEvent event = new GenericWrapperEvent<>(new TestEvent("event-id", "event-payload")); + sqsTemplate.send(INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE, event); + logger.debug("Sent message GenericWrapperEvent"); + + assertThat(latchContainer.infersGenericWrapperPayloadLatch.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(pojoCollector.receivedGenericWrapperPayload).hasSize(1); + GenericWrapperEvent receivedPayload = pojoCollector.receivedGenericWrapperPayload.get(0); + Object genericPayload = receivedPayload.testEvent(); + assertThat(genericPayload).isInstanceOf(TestEvent.class); + assertThat(receivedPayload).isEqualTo(event); + interceptorPayloadTypeCollector.assertPayloadsForQueueContains(INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE, event); + assertThat(ackLatch.await(10, TimeUnit.SECONDS)).isTrue(); + ackCallbackPayloadTypeCollector.assertPayloadsForQueueContains(INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE, event); + + } + + @Test + void shouldInferTestEventTypeFromBatchOfGenericWrappers() throws Exception { + CountDownLatch ackLatch = new CountDownLatch(2); + ackCallbackPayloadTypeCollector.registerLatch(INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE, ackLatch); + List> events = List.of( + new GenericWrapperEvent<>(new TestEvent("batch-wrapper-id-1", "batch-wrapper-payload-1")), + new GenericWrapperEvent<>(new TestEvent("batch-wrapper-id-2", "batch-wrapper-payload-2"))); + + sqsTemplate.sendMany(INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE, + events.stream().map(event -> MessageBuilder.withPayload(event).build()).toList()); + + assertThat(latchContainer.infersBatchGenericWrapperPayloadLatch.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(pojoCollector.receivedBatchGenericWrapperPayload).containsExactlyInAnyOrderElementsOf(events) + .allSatisfy(wrapper -> assertThat(wrapper.testEvent()).isInstanceOf(TestEvent.class)); + interceptorPayloadTypeCollector.assertPayloadsForQueueContainsAll(INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE, + events); + assertThat(ackLatch.await(10, TimeUnit.SECONDS)).isTrue(); + ackCallbackPayloadTypeCollector.assertPayloadsForQueueContainsAll(INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE, + events); + } + + @Test + void shouldInferTestEventTypeFromBatchOfMessagesWithGenericWrappers() throws Exception { + CountDownLatch ackLatch = new CountDownLatch(2); + ackCallbackPayloadTypeCollector.registerLatch(INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE, ackLatch); + List> events = List.of( + new GenericWrapperEvent<>( + new TestEvent("batch-message-wrapper-id-1", "batch-message-wrapper-payload-1")), + new GenericWrapperEvent<>( + new TestEvent("batch-message-wrapper-id-2", "batch-message-wrapper-payload-2"))); + + sqsTemplate.sendMany(INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE, + events.stream().map(event -> MessageBuilder.withPayload(event).build()).toList()); + + assertThat(latchContainer.infersBatchMessageGenericWrapperPayloadLatch.await(10, TimeUnit.SECONDS)).isTrue(); + assertThat(pojoCollector.receivedBatchMessageGenericWrapperPayload).containsExactlyInAnyOrderElementsOf(events) + .allSatisfy(wrapper -> assertThat(wrapper.testEvent()).isInstanceOf(TestEvent.class)); + interceptorPayloadTypeCollector + .assertPayloadsForQueueContainsAll(INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE, events); + assertThat(ackLatch.await(10, TimeUnit.SECONDS)).isTrue(); + ackCallbackPayloadTypeCollector + .assertPayloadsForQueueContainsAll(INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE, events); + } + static class InfersSimplePojoListener { @Autowired @@ -405,6 +503,72 @@ void listen(Message message) { } + static class InfersMessageListPayloadListener { + + @Autowired + LatchContainer latchContainer; + + @Autowired + PojoCollector pojoCollector; + + @SqsListener(queueNames = INFERS_MESSAGE_LIST_PAYLOAD_QUEUE, id = "infers-message-list-payload") + void listen(Message> message) { + logger.debug("Received message with List payload: {}", message); + pojoCollector.receivedMessageListPayloads.add(message); + latchContainer.infersMessageListPayloadLatch.countDown(); + } + + } + + static class InfersGenericWrapperPayloadListener { + @Autowired + LatchContainer latchContainer; + + @Autowired + PojoCollector pojoCollector; + + @SqsListener(queueNames = INFERS_GENERIC_WRAPPER_PAYLOAD_QUEUE, id = "infers-generic-wrapper-payload") + void listen(GenericWrapperEvent message) { + logger.debug("Received message with GenericWrapperEvent payload: {}", message); + pojoCollector.receivedGenericWrapperPayload.add(message); + latchContainer.infersGenericWrapperPayloadLatch.countDown(); + } + } + + static class InfersBatchGenericWrapperPayloadListener { + + @Autowired + LatchContainer latchContainer; + + @Autowired + PojoCollector pojoCollector; + + @SqsListener(queueNames = INFERS_BATCH_GENERIC_WRAPPER_PAYLOAD_QUEUE, id = "infers-batch-generic-wrapper-payload") + void listen(List> messages) { + logger.debug("Received {} GenericWrapperEvent payloads", messages.size()); + pojoCollector.receivedBatchGenericWrapperPayload.addAll(messages); + messages.forEach(message -> latchContainer.infersBatchGenericWrapperPayloadLatch.countDown()); + } + } + + static class InfersBatchMessageGenericWrapperPayloadListener { + + @Autowired + LatchContainer latchContainer; + + @Autowired + PojoCollector pojoCollector; + + @SqsListener(queueNames = INFERS_BATCH_MESSAGE_GENERIC_WRAPPER_PAYLOAD_QUEUE, id = "infers-batch-message-generic-wrapper-payload") + void listen(List>> messages) { + logger.debug("Received {} messages with GenericWrapperEvent payloads", messages.size()); + messages.forEach(message -> { + pojoCollector.receivedBatchMessageGenericWrapperPayload.add(message.getPayload()); + latchContainer.infersBatchMessageGenericWrapperPayloadLatch.countDown(); + }); + } + } + static class AsyncListener { @Autowired @@ -544,6 +708,17 @@ static class PojoCollector { final List receivedCustomConverterPojos = Collections.synchronizedList(new ArrayList<>()); + final List>> receivedMessageListPayloads = Collections + .synchronizedList(new ArrayList<>()); + + final List> receivedGenericWrapperPayload = Collections + .synchronizedList(new ArrayList<>()); + + final List> receivedBatchGenericWrapperPayload = Collections + .synchronizedList(new ArrayList<>()); + + final List> receivedBatchMessageGenericWrapperPayload = Collections + .synchronizedList(new ArrayList<>()); } /** @@ -658,6 +833,13 @@ static class LatchContainer { final CountDownLatch customConverterLatch = new CountDownLatch(1); + final CountDownLatch infersMessageListPayloadLatch = new CountDownLatch(1); + + final CountDownLatch infersGenericWrapperPayloadLatch = new CountDownLatch(1); + + final CountDownLatch infersBatchGenericWrapperPayloadLatch = new CountDownLatch(2); + + final CountDownLatch infersBatchMessageGenericWrapperPayloadLatch = new CountDownLatch(2); } @Import(SqsBootstrapConfiguration.class) @@ -885,6 +1067,26 @@ SqsTemplate sqsTemplate() { .configureDefaultConverter(AbstractMessagingMessageConverter::doNotSendPayloadTypeHeader).build(); } + @Bean + InfersMessageListPayloadListener infersMessageListPayloadListener() { + return new InfersMessageListPayloadListener(); + } + + @Bean + InfersGenericWrapperPayloadListener infersGenericWrapperPayloadListener() { + return new InfersGenericWrapperPayloadListener(); + } + + @Bean + InfersBatchGenericWrapperPayloadListener infersBatchGenericWrapperPayloadListener() { + return new InfersBatchGenericWrapperPayloadListener(); + } + + @Bean + InfersBatchMessageGenericWrapperPayloadListener infersBatchMessageGenericWrapperPayloadListener() { + return new InfersBatchMessageGenericWrapperPayloadListener(); + } + } static class TestEvent { @@ -1043,4 +1245,8 @@ public String toString() { } + record GenericWrapperEvent(T testEvent) + { + } + } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainerTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainerTests.java index f7e450286..b3d26b014 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainerTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/AbstractMessageListenerContainerTests.java @@ -106,4 +106,53 @@ void shouldSetAsyncComponents() { } + @Test + void shouldDelegatePayloadTypeOnlySetterToConversionHintAwareOverload() { + RecordingMessageListenerContainer container = new RecordingMessageListenerContainer( + SqsContainerOptions.builder().build()); + + container.setPayloadDeserializationType(String.class); + + assertThat(container.payloadType).isEqualTo(String.class); + assertThat(container.conversionHint).isNull(); + assertThat(container.conversionHintAwareSetterInvocations).isEqualTo(1); + } + + @Test + void shouldClearConversionHintWhenPayloadTypeIsCleared() { + AbstractMessageListenerContainer container = new AbstractMessageListenerContainer<>( + SqsContainerOptions.builder().build()) { + }; + Object conversionHint = new Object(); + container.setPayloadDeserializationType(String.class, conversionHint); + + container.setPayloadDeserializationType(null, conversionHint); + + assertThat(container.getPayloadDeserializationType()).isNull(); + assertThat(container.getPayloadConversionHint()).isNull(); + } + + private static class RecordingMessageListenerContainer + extends AbstractMessageListenerContainer { + + private int conversionHintAwareSetterInvocations; + + private Class payloadType; + + private Object conversionHint; + + RecordingMessageListenerContainer(SqsContainerOptions containerOptions) { + super(containerOptions); + } + + @Override + public void setPayloadDeserializationType(Class payloadType, Object conversionHint) { + this.conversionHintAwareSetterInvocations++; + this.payloadType = payloadType; + this.conversionHint = conversionHint; + super.setPayloadDeserializationType(payloadType, conversionHint); + } + + } + } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/SqsMessageListenerContainerTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/SqsMessageListenerContainerTests.java index 3489ac5d3..3374cfc1c 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/SqsMessageListenerContainerTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/SqsMessageListenerContainerTests.java @@ -32,10 +32,13 @@ import io.awspring.cloud.sqs.listener.interceptor.MessageInterceptor; import io.awspring.cloud.sqs.listener.pipeline.MessageProcessingPipeline; import io.awspring.cloud.sqs.listener.sink.MessageSink; +import io.awspring.cloud.sqs.listener.source.AbstractMessageConvertingMessageSource; import io.awspring.cloud.sqs.listener.source.MessageSource; import io.awspring.cloud.sqs.support.observation.AbstractListenerObservation; import io.awspring.cloud.sqs.support.observation.SqsListenerObservation; import io.micrometer.observation.ObservationRegistry; +import java.lang.reflect.Field; +import java.lang.reflect.Method; import java.util.Arrays; import java.util.Collections; import java.util.List; @@ -43,6 +46,7 @@ import java.util.concurrent.CompletionException; import org.junit.jupiter.api.Test; import org.mockito.Mockito; +import org.springframework.core.MethodParameter; import org.springframework.core.task.SimpleAsyncTaskExecutor; import software.amazon.awssdk.services.sqs.SqsAsyncClient; import software.amazon.awssdk.services.sqs.model.GetQueueUrlRequest; @@ -127,6 +131,36 @@ void shouldCreateFromBuilderWithAsyncComponents() { assertThat(container.getPhase()).isEqualTo(MessageListenerContainer.DEFAULT_PHASE); } + @Test + void shouldPropagatePayloadClassAndConversionHintToMessageSource() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + RecordingMessageSource messageSource = new RecordingMessageSource(); + TestSqsMessageListenerContainer container = new TestSqsMessageListenerContainer(mock(SqsAsyncClient.class), + SqsContainerOptions.builder().build()); + container.setId("test-container"); + container.setPayloadDeserializationType(GenericWrapper.class, conversionHint); + + container.configureMessageSource(messageSource); + + assertThat(messageSource.payloadClass).isEqualTo(GenericWrapper.class); + assertThat(messageSource.conversionHint).isSameAs(conversionHint); + } + + @Test + void shouldClearConversionHintWhenPayloadDeserializationTypeIsSetWithoutHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + TestSqsMessageListenerContainer container = new TestSqsMessageListenerContainer(mock(SqsAsyncClient.class), + SqsContainerOptions.builder().build()); + container.setPayloadDeserializationType(GenericWrapper.class, conversionHint); + + container.setPayloadDeserializationType(String.class); + + assertThat(container.getPayloadDeserializationType()).isEqualTo(String.class); + assertThat(container.getPayloadConversionHint()).isNull(); + } + @Test void shouldThrowIfWrongCustomExecutor() { SqsAsyncClient client = mock(SqsAsyncClient.class); @@ -308,10 +342,31 @@ public MessageSink getCreatedMessageSink() { return this.createdMessageSink; } + void configureMessageSource(MessageSource messageSource) { + try { + Field field = AbstractPipelineMessageListenerContainer.class.getDeclaredField("messageSources"); + field.setAccessible(true); + field.set(this, List.of(messageSource)); + configureMessageSources(new ContainerComponentFactory<>() { + @Override + public MessageSource createMessageSource(SqsContainerOptions options) { + return messageSource; + } + + @Override + public MessageSink createMessageSink(SqsContainerOptions options) { + return (messages, context) -> CompletableFuture.completedFuture(null); + } + }); + } + catch (ReflectiveOperationException e) { + throw new IllegalStateException("Could not configure test message source", e); + } + } + private MessageSink getMessageSink() { try { - java.lang.reflect.Field field = AbstractPipelineMessageListenerContainer.class - .getDeclaredField("messageSink"); + Field field = AbstractPipelineMessageListenerContainer.class.getDeclaredField("messageSink"); field.setAccessible(true); return (MessageSink) field.get(this); } @@ -321,4 +376,37 @@ private MessageSink getMessageSink() { } } + private static class RecordingMessageSource extends AbstractMessageConvertingMessageSource { + + private Class payloadClass; + + private Object conversionHint; + + @Override + public void setPayloadDeserializationType(Class payloadDeserializationType) { + this.payloadClass = payloadDeserializationType; + } + + @Override + public void setPayloadDeserializationType(Class payloadDeserializationType, Object conversionHint) { + this.payloadClass = payloadDeserializationType; + this.conversionHint = conversionHint; + } + + @Override + public void setMessageSink(MessageSink messageSink) { + } + + } + + static class GenericListener { + + void listen(GenericWrapper payload) { + } + + } + + static class GenericWrapper { + } + } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSourceTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSourceTests.java new file mode 100644 index 000000000..d901445e2 --- /dev/null +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractMessageConvertingMessageSourceTests.java @@ -0,0 +1,96 @@ +/* + * Copyright 2013-2026 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. + * You may obtain a copy of the License at + * + * https://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 io.awspring.cloud.sqs.listener.source; + +import static org.assertj.core.api.Assertions.assertThat; + +import io.awspring.cloud.sqs.listener.SqsContainerOptions; +import io.awspring.cloud.sqs.listener.sink.MessageSink; +import io.awspring.cloud.sqs.support.converter.MessageConversionContext; +import org.junit.jupiter.api.Test; + +/** + * Tests for {@link AbstractMessageConvertingMessageSource}. + * + * @author Bruno Augusto Garcia + */ +class AbstractMessageConvertingMessageSourceTests { + + @Test + void shouldDelegatePayloadTypeOnlySetterToConversionHintAwareHook() { + ConversionHintHookRecordingMessageSource source = new ConversionHintHookRecordingMessageSource(); + source.configure(SqsContainerOptions.builder().build()); + + source.setPayloadDeserializationType(String.class); + + assertThat(source.payloadType).isEqualTo(String.class); + assertThat(source.conversionHint).isNull(); + assertThat(source.context).isSameAs(source.getMessageConversionContext()); + } + + @Test + void shouldInvokeLegacyHookFromConversionHintAwareSetter() { + LegacyHookRecordingMessageSource source = new LegacyHookRecordingMessageSource(); + source.configure(SqsContainerOptions.builder().build()); + Object conversionHint = new Object(); + + source.setPayloadDeserializationType(String.class, conversionHint); + + assertThat(source.payloadType).isEqualTo(String.class); + assertThat(source.context).isSameAs(source.getMessageConversionContext()); + } + + private abstract static class TestMessageSource extends AbstractMessageConvertingMessageSource { + + @Override + public void setMessageSink(MessageSink messageSink) { + } + + } + + private static class ConversionHintHookRecordingMessageSource extends TestMessageSource { + + private Class payloadType; + + private Object conversionHint; + + private MessageConversionContext context; + + @Override + protected void doConfigurePayloadTypeOnContext(Class payloadType, Object conversionHint, + MessageConversionContext context) { + this.payloadType = payloadType; + this.conversionHint = conversionHint; + this.context = context; + } + + } + + private static class LegacyHookRecordingMessageSource extends TestMessageSource { + + private Class payloadType; + + private MessageConversionContext context; + + @Override + protected void doConfigurePayloadTypeOnContext(Class payloadType, MessageConversionContext context) { + this.payloadType = payloadType; + this.context = context; + } + + } + +} diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSourceTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSourceTests.java index 95c3aa215..72e67b1fd 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSourceTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/listener/source/AbstractSqsMessageSourceTests.java @@ -22,6 +22,10 @@ import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; +import io.awspring.cloud.sqs.listener.SqsContainerOptions; +import io.awspring.cloud.sqs.support.converter.MessageConversionContext; +import io.awspring.cloud.sqs.support.converter.SqsMessageConversionContext; +import java.lang.reflect.Method; import java.util.Collection; import java.util.List; import java.util.UUID; @@ -30,6 +34,7 @@ import java.util.stream.IntStream; import org.junit.jupiter.api.Test; import org.mockito.ArgumentCaptor; +import org.springframework.core.MethodParameter; import software.amazon.awssdk.services.sqs.SqsAsyncClient; import software.amazon.awssdk.services.sqs.model.Message; import software.amazon.awssdk.services.sqs.model.ReceiveMessageRequest; @@ -92,8 +97,79 @@ void shouldRequestHundredAndOneMessages() { assertThat(requests.get(10)).extracting(ReceiveMessageRequest::maxNumberOfMessages).isEqualTo(1); } + @Test + void shouldConfigurePayloadClassAndConversionHintOnContext() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + AbstractSqsMessageSource source = new StandardSqsMessageSource<>(); + source.configure(SqsContainerOptions.builder().build()); + + source.setPayloadDeserializationType(GenericWrapper.class, conversionHint); + + assertThat(source.getMessageConversionContext()).isInstanceOfSatisfying(SqsMessageConversionContext.class, + context -> { + assertThat(context.getPayloadClass()).isEqualTo(GenericWrapper.class); + assertThat(context.getConversionHint()).isSameAs(conversionHint); + }); + } + + @Test + void shouldInvokeLegacyPayloadConfigurationHookFromConversionHintAwareSetter() { + LegacyHookRecordingSqsMessageSource source = new LegacyHookRecordingSqsMessageSource<>(); + source.configure(SqsContainerOptions.builder().build()); + Object conversionHint = new Object(); + + source.setPayloadDeserializationType(String.class, conversionHint); + + assertThat(source.legacyHookInvoked).isTrue(); + assertThat(source.getMessageConversionContext()).isInstanceOfSatisfying(SqsMessageConversionContext.class, + context -> { + assertThat(context.getPayloadClass()).isEqualTo(String.class); + assertThat(context.getConversionHint()).isSameAs(conversionHint); + }); + } + + @Test + void shouldClearConversionHintWhenPayloadTypeIsReconfiguredThroughLegacySetter() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + AbstractSqsMessageSource source = new StandardSqsMessageSource<>(); + source.configure(SqsContainerOptions.builder().build()); + source.setPayloadDeserializationType(GenericWrapper.class, conversionHint); + + source.setPayloadDeserializationType(String.class); + + assertThat(source.getMessageConversionContext()).isInstanceOfSatisfying(SqsMessageConversionContext.class, + context -> { + assertThat(context.getPayloadClass()).isEqualTo(String.class); + assertThat(context.getConversionHint()).isNull(); + }); + } + private List getHundredMessages(List batch) { return IntStream.range(0, 10).mapToObj(index -> batch).flatMap(Collection::stream).collect(Collectors.toList()); } + static class GenericListener { + + void listen(GenericWrapper payload) { + } + + } + + static class GenericWrapper { + } + + private static class LegacyHookRecordingSqsMessageSource extends StandardSqsMessageSource { + + private boolean legacyHookInvoked; + + @Override + protected void doConfigurePayloadTypeOnContext(Class payloadType, MessageConversionContext context) { + this.legacyHookInvoked = true; + super.doConfigurePayloadTypeOnContext(payloadType, context); + } + + } + } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/LegacyJackson2SqsMessagingMessageConverterTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/LegacyJackson2SqsMessagingMessageConverterTests.java index 549e190ef..a7ec9ba04 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/LegacyJackson2SqsMessagingMessageConverterTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/LegacyJackson2SqsMessagingMessageConverterTests.java @@ -24,10 +24,13 @@ import com.fasterxml.jackson.databind.ObjectMapper; import io.awspring.cloud.sqs.support.converter.legacy.LegacyJackson2SqsMessagingMessageConverter; +import java.lang.reflect.Method; import java.util.Collections; +import java.util.List; import java.util.Objects; import java.util.UUID; import org.junit.jupiter.api.Test; +import org.springframework.core.MethodParameter; import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.converter.MessageConverter; import org.springframework.messaging.support.MessageBuilder; @@ -156,10 +159,79 @@ void shouldReturnFalseAfterSettingPayloadTypeHeader() { assertThat(converter.isUsingDefaultPayloadTypeMapper()).isFalse(); } + @Test + void shouldDeserializeGenericWrapperUsingListenerMethodConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + LegacyJackson2SqsMessagingMessageConverter converter = new LegacyJackson2SqsMessagingMessageConverter(); + SqsMessageConversionContext context = new SqsMessageConversionContext(); + context.setPayloadClass(GenericWrapper.class); + context.setConversionHint(conversionHint); + Message message = Message.builder().body(""" + {"value":{"myProperty":"nested-value"}} + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isInstanceOfSatisfying(GenericWrapper.class, + wrapper -> assertThat(wrapper.getValue()).isInstanceOfSatisfying(MyPojo.class, + pojo -> assertThat(pojo.getMyProperty()).isEqualTo("nested-value"))); + } + + @Test + void shouldDeserializeNestedGenericWrapperUsingListenerMethodConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listenToNestedWrapper", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + LegacyJackson2SqsMessagingMessageConverter converter = new LegacyJackson2SqsMessagingMessageConverter(); + SqsMessageConversionContext context = new SqsMessageConversionContext(); + context.setPayloadClass(GenericWrapper.class); + context.setConversionHint(conversionHint); + Message message = Message.builder().body(""" + {"value":[{"myProperty":"first"},{"myProperty":"second"}]} + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isInstanceOfSatisfying(GenericWrapper.class, + wrapper -> assertThat(wrapper.getValue()) + .isEqualTo(List.of(new MyPojo("first"), new MyPojo("second")))); + } + + static class GenericListener { + + void listen(GenericWrapper payload) { + } + + void listenToNestedWrapper(GenericWrapper> payload) { + } + + } + + static class GenericWrapper { + + private T value; + + public T getValue() { + return this.value; + } + + public void setValue(T value) { + this.value = value; + } + + } + static class MyPojo { private String myProperty = "myValue"; + MyPojo() { + } + + MyPojo(String myProperty) { + this.myProperty = myProperty; + } + public String getMyProperty() { return this.myProperty; } diff --git a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/SqsMessagingMessageConverterTests.java b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/SqsMessagingMessageConverterTests.java index 1864cca97..3e01ee5d6 100644 --- a/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/SqsMessagingMessageConverterTests.java +++ b/spring-cloud-aws-sqs/src/test/java/io/awspring/cloud/sqs/support/converter/SqsMessagingMessageConverterTests.java @@ -19,11 +19,17 @@ import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.assertj.core.api.InstanceOfAssertFactories.type; +import java.lang.reflect.Method; +import java.util.List; import java.util.Objects; import java.util.UUID; import org.junit.jupiter.api.Test; +import org.springframework.core.MethodParameter; +import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.converter.CompositeMessageConverter; import org.springframework.messaging.converter.JacksonJsonMessageConverter; +import org.springframework.messaging.converter.MessageConverter; +import org.springframework.messaging.converter.SmartMessageConverter; import software.amazon.awssdk.services.sqs.model.Message; import tools.jackson.databind.json.JsonMapper; @@ -77,6 +83,206 @@ void shouldConvertMessageWithCustomJsonMapper() throws Exception { assertThat(resultMessage.getPayload()).isEqualTo(myPojo); } + @Test + void shouldPassConversionHintToSmartPayloadConverter() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + RecordingSmartMessageConverter payloadConverter = new RecordingSmartMessageConverter(); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + converter.setPayloadMessageConverter(payloadConverter); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body("{}").messageId(UUID.randomUUID().toString()).build(); + + converter.toMessagingMessage(message, context); + + assertThat(payloadConverter.targetClass).isEqualTo(GenericWrapper.class); + assertThat(payloadConverter.conversionHint).isSameAs(conversionHint); + assertThat(payloadConverter.smartOverloadInvoked).isTrue(); + } + + @Test + void shouldNotUseInferredConversionHintWhenCustomPayloadTypeMapperTakesPrecedence() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + RecordingSmartMessageConverter payloadConverter = new RecordingSmartMessageConverter(); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + converter.setPayloadMessageConverter(payloadConverter); + converter.setPayloadTypeMapper(message -> String.class); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body("{}").messageId(UUID.randomUUID().toString()).build(); + + converter.toMessagingMessage(message, context); + + assertThat(payloadConverter.targetClass).isEqualTo(String.class); + assertThat(payloadConverter.conversionHint).isNull(); + assertThat(payloadConverter.smartOverloadInvoked).isFalse(); + } + + @Test + void shouldFallBackToRegularPayloadConverterWhenConversionHintIsPresent() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + RecordingMessageConverter payloadConverter = new RecordingMessageConverter(); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + converter.setPayloadMessageConverter(payloadConverter); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body("{}").messageId(UUID.randomUUID().toString()).build(); + + converter.toMessagingMessage(message, context); + + assertThat(payloadConverter.targetClass).isEqualTo(GenericWrapper.class); + } + + @Test + void shouldDeserializeGenericWrapperUsingListenerMethodConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listen", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body(""" + {"value":{"name":"nested-value"}} + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isInstanceOfSatisfying(GenericWrapper.class, + wrapper -> assertThat(wrapper.value()).isEqualTo(new NestedPojo("nested-value"))); + } + + @Test + void shouldDeserializeMessageWithGenericCollectionPayloadUsingListenerMethodConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listenToMessage", + org.springframework.messaging.Message.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + SqsMessageConversionContext context = createConversionContext(List.class, conversionHint); + Message message = Message.builder().body(""" + [{"name":"first"},{"name":"second"}] + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isEqualTo(List.of(new NestedPojo("first"), new NestedPojo("second"))); + } + + @Test + void shouldDeserializeGenericWrapperFromBatchElementConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listenToBatch", List.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0).nested(); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body(""" + {"value":{"name":"batch-value"}} + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isInstanceOfSatisfying(GenericWrapper.class, + wrapper -> assertThat(wrapper.value()).isEqualTo(new NestedPojo("batch-value"))); + } + + @Test + void shouldDeserializeNestedGenericWrapperUsingListenerMethodConversionHint() throws Exception { + Method listenerMethod = GenericListener.class.getDeclaredMethod("listenToNestedWrapper", GenericWrapper.class); + MethodParameter conversionHint = new MethodParameter(listenerMethod, 0); + SqsMessagingMessageConverter converter = new SqsMessagingMessageConverter(); + SqsMessageConversionContext context = createConversionContext(GenericWrapper.class, conversionHint); + Message message = Message.builder().body(""" + {"value":[{"name":"first"},{"name":"second"}]} + """).messageId(UUID.randomUUID().toString()).build(); + + org.springframework.messaging.Message result = converter.toMessagingMessage(message, context); + + assertThat(result.getPayload()).isInstanceOfSatisfying(GenericWrapper.class, + wrapper -> assertThat(wrapper.value()) + .isEqualTo(List.of(new NestedPojo("first"), new NestedPojo("second")))); + } + + private SqsMessageConversionContext createConversionContext(Class payloadClass, Object conversionHint) { + SqsMessageConversionContext context = new SqsMessageConversionContext(); + context.setPayloadClass(payloadClass); + context.setConversionHint(conversionHint); + return context; + } + + static class GenericListener { + + void listen(GenericWrapper payload) { + } + + void listenToMessage(org.springframework.messaging.Message> message) { + } + + void listenToBatch(List> payload) { + } + + void listenToNestedWrapper(GenericWrapper> payload) { + } + + } + + record GenericWrapper(T value) + { + } + + record NestedPojo(String name) { + } + + static class RecordingSmartMessageConverter implements SmartMessageConverter { + + private Class targetClass; + + private Object conversionHint; + + private boolean smartOverloadInvoked; + + @Override + public Object fromMessage(org.springframework.messaging.Message message, Class targetClass) { + this.targetClass = targetClass; + this.conversionHint = null; + this.smartOverloadInvoked = false; + return "converted"; + } + + @Override + public Object fromMessage(org.springframework.messaging.Message message, Class targetClass, + Object conversionHint) { + this.targetClass = targetClass; + this.conversionHint = conversionHint; + this.smartOverloadInvoked = true; + return "converted"; + } + + @Override + public org.springframework.messaging.Message toMessage(Object payload, MessageHeaders headers) { + return null; + } + + @Override + public org.springframework.messaging.Message toMessage(Object payload, MessageHeaders headers, + Object conversionHint) { + return null; + } + + } + + static class RecordingMessageConverter implements MessageConverter { + + private Class targetClass; + + @Override + public Object fromMessage(org.springframework.messaging.Message message, Class targetClass) { + this.targetClass = targetClass; + return "converted"; + } + + @Override + public org.springframework.messaging.Message toMessage(Object payload, MessageHeaders headers) { + return null; + } + + } + static class MyPojo { private String myProperty = "myValue"; @@ -105,4 +311,4 @@ public int hashCode() { } } -} \ No newline at end of file +}