diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/DefaultStompSession.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/DefaultStompSession.java index c2512edc03..3bd04271ff 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/DefaultStompSession.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/stomp/DefaultStompSession.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 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. @@ -256,12 +256,9 @@ public class DefaultStompSession implements ConnectionHandlingStompSession { private Message createMessage(StompHeaderAccessor accessor, @Nullable Object payload) { accessor.updateSimpMessageHeadersFromStompHeaders(); Message message; - if (payload == null) { + if (isEmpty(payload)) { message = MessageBuilder.createMessage(EMPTY_PAYLOAD, accessor.getMessageHeaders()); } - else if (payload instanceof byte[]) { - message = MessageBuilder.createMessage((byte[]) payload, accessor.getMessageHeaders()); - } else { message = (Message) getMessageConverter().toMessage(payload, accessor.getMessageHeaders()); accessor.updateStompHeadersFromSimpMessageHeaders(); @@ -274,6 +271,11 @@ public class DefaultStompSession implements ConnectionHandlingStompSession { return message; } + private boolean isEmpty(@Nullable Object payload) { + return payload == null || StringUtils.isEmpty(payload) || + (payload instanceof byte[] && ((byte[]) payload).length == 0); + } + private void execute(Message message) { if (logger.isTraceEnabled()) { StompHeaderAccessor accessor = MessageHeaderAccessor.getAccessor(message, StompHeaderAccessor.class); diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/DefaultStompSessionTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/DefaultStompSessionTests.java index 58c9355254..451a2ef11b 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/DefaultStompSessionTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/stomp/DefaultStompSessionTests.java @@ -17,6 +17,7 @@ package org.springframework.messaging.simp.stomp; import java.nio.charset.StandardCharsets; +import java.util.Arrays; import java.util.Date; import java.util.Map; import java.util.concurrent.ScheduledFuture; @@ -32,6 +33,8 @@ import org.mockito.junit.MockitoJUnitRunner; import org.springframework.messaging.Message; import org.springframework.messaging.MessageDeliveryException; +import org.springframework.messaging.converter.ByteArrayMessageConverter; +import org.springframework.messaging.converter.CompositeMessageConverter; import org.springframework.messaging.converter.MessageConversionException; import org.springframework.messaging.converter.StringMessageConverter; import org.springframework.messaging.simp.stomp.StompSession.Receiptable; @@ -83,7 +86,9 @@ public class DefaultStompSessionTests { public void setUp() { this.connectHeaders = new StompHeaders(); this.session = new DefaultStompSession(this.sessionHandler, this.connectHeaders); - this.session.setMessageConverter(new StringMessageConverter()); + this.session.setMessageConverter( + new CompositeMessageConverter( + Arrays.asList(new StringMessageConverter(), new ByteArrayMessageConverter()))); SettableListenableFuture future = new SettableListenableFuture<>(); future.set(null); @@ -111,7 +116,7 @@ public class DefaultStompSessionTests { @Test // SPR-16844 public void afterConnectedWithSpecificVersion() { assertThat(this.session.isConnected()).isFalse(); - this.connectHeaders.setAcceptVersion(new String[] {"1.1"}); + this.connectHeaders.setAcceptVersion("1.1"); this.session.afterConnected(this.connection); @@ -390,6 +395,26 @@ public class DefaultStompSessionTests { assertThat(accessor.getReceipt()).isEqualTo("my-receipt"); } + @Test // gh-23358 + public void sendByteArray() { + this.session.afterConnected(this.connection); + assertThat(this.session.isConnected()); + + String destination = "/topic/foo"; + String payload = "sample payload"; + this.session.send(destination, payload.getBytes(StandardCharsets.UTF_8)); + + Message message = this.messageCaptor.getValue(); + StompHeaderAccessor accessor = MessageHeaderAccessor.getAccessor(message, StompHeaderAccessor.class); + + StompHeaders stompHeaders = StompHeaders.readOnlyStompHeaders(accessor.getNativeHeaders()); + assertThat(stompHeaders.size()).as(stompHeaders.toString()).isEqualTo(2); + + assertThat(stompHeaders.getDestination()).isEqualTo(destination); + assertThat(stompHeaders.getContentType()).isEqualTo(MimeTypeUtils.APPLICATION_OCTET_STREAM); + assertThat(new String(message.getPayload(), StandardCharsets.UTF_8)).isEqualTo(payload); + } + @Test public void sendWithConversionException() { this.session.afterConnected(this.connection);