Add BufferingStompDecoder

Before this change the StompDecoder decoded and returned only the first
Message in the ByteBuffer passed to it. So to obtain all messages from
the buffer, one had to loop passing the same buffer in until no more
complete STOMP frames could be decoded.

This chage modifies StompDecoder to return List<Message> after
exhaustively decoding all available STOMP frames from the input buffer.
Also an overloaded decode method allows passing in Map that will be
populated with any headers successfully parsed, which is useful for
"peeking" at the "content-length" header.

This change also adds a BufferingStompDecoder sub-class which buffers
any content left in the input buffer after parsing one or more STOMP
frames. This sub-class can also deal with fragmented messages,
re-assembling them and parsing as a whole message.

Issue: SPR-11527
This commit is contained in:
Rossen Stoyanchev
2014-03-21 11:32:39 -04:00
parent 465ca24ab2
commit ebffd67b5e
7 changed files with 482 additions and 74 deletions

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2013 the original author or authors.
* Copyright 2002-2014 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.
@@ -21,7 +21,9 @@ import java.nio.ByteBuffer;
import java.security.Principal;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
@@ -29,6 +31,7 @@ import org.apache.commons.logging.LogFactory;
import org.springframework.messaging.Message;
import org.springframework.messaging.MessageChannel;
import org.springframework.messaging.simp.SimpMessageType;
import org.springframework.messaging.simp.stomp.BufferingStompDecoder;
import org.springframework.messaging.simp.stomp.StompCommand;
import org.springframework.messaging.simp.stomp.StompConversionException;
import org.springframework.messaging.simp.stomp.StompDecoder;
@@ -66,13 +69,31 @@ public class StompSubProtocolHandler implements SubProtocolHandler {
private static final Log logger = LogFactory.getLog(StompSubProtocolHandler.class);
private final StompDecoder stompDecoder = new StompDecoder();
private int messageBufferSizeLimit = 64 * 1024;
private final Map<String, BufferingStompDecoder> decoders = new ConcurrentHashMap<String, BufferingStompDecoder>();
private final StompEncoder stompEncoder = new StompEncoder();
private UserSessionRegistry userSessionRegistry;
/**
* TODO
* @param messageBufferSizeLimit
*/
public void setMessageBufferSizeLimit(int messageBufferSizeLimit) {
this.messageBufferSizeLimit = messageBufferSizeLimit;
}
/**
* TODO
* @return
*/
public int getMessageBufferSizeLimit() {
return this.messageBufferSizeLimit;
}
/**
* Provide a registry with which to register active user session ids.
* @see org.springframework.messaging.simp.user.UserDestinationMessageHandler
@@ -99,49 +120,53 @@ public class StompSubProtocolHandler implements SubProtocolHandler {
public void handleMessageFromClient(WebSocketSession session,
WebSocketMessage<?> webSocketMessage, MessageChannel outputChannel) {
Message<?> message = null;
Throwable decodeFailure = null;
List<Message<byte[]>> messages = null;
try {
Assert.isInstanceOf(TextMessage.class, webSocketMessage);
TextMessage textMessage = (TextMessage) webSocketMessage;
ByteBuffer byteBuffer = ByteBuffer.wrap(textMessage.asBytes());
message = this.stompDecoder.decode(byteBuffer);
if (message == null) {
decodeFailure = new IllegalStateException("Not a valid STOMP frame: " + textMessage.getPayload());
BufferingStompDecoder decoder = this.decoders.get(session.getId());
if (decoder == null) {
throw new IllegalStateException("No decoder for session id '" + session.getId() + "'");
}
messages = decoder.decode(byteBuffer);
if (messages.isEmpty()) {
logger.debug("Incomplete STOMP frame content received," + "buffered=" +
decoder.getBufferSize() + ", buffer size limit=" + decoder.getBufferSizeLimit());
return;
}
}
catch (Throwable ex) {
decodeFailure = ex;
}
if (decodeFailure != null) {
logger.error("Failed to parse WebSocket message as STOMP frame", decodeFailure);
sendErrorMessage(session, decodeFailure);
logger.error("Failed to parse WebSocket message to STOMP frame(s)", ex);
sendErrorMessage(session, ex);
return;
}
try {
StompHeaderAccessor headers = StompHeaderAccessor.wrap(message);
if (logger.isTraceEnabled()) {
if (SimpMessageType.HEARTBEAT.equals(headers.getMessageType())) {
logger.trace("Received heartbeat from client session=" + session.getId());
}
else {
logger.trace("Received message from client session=" + session.getId());
for (Message<byte[]> message : messages) {
try {
StompHeaderAccessor headers = StompHeaderAccessor.wrap(message);
if (logger.isTraceEnabled()) {
if (SimpMessageType.HEARTBEAT.equals(headers.getMessageType())) {
logger.trace("Received heartbeat from client session=" + session.getId());
}
else {
logger.trace("Received message from client session=" + session.getId());
}
}
headers.setSessionId(session.getId());
headers.setSessionAttributes(session.getAttributes());
headers.setUser(session.getPrincipal());
message = MessageBuilder.withPayload(message.getPayload()).setHeaders(headers).build();
outputChannel.send(message);
}
catch (Throwable ex) {
logger.error("Terminating STOMP session due to failure to send message", ex);
sendErrorMessage(session, ex);
}
headers.setSessionId(session.getId());
headers.setSessionAttributes(session.getAttributes());
headers.setUser(session.getPrincipal());
message = MessageBuilder.withPayload(message.getPayload()).setHeaders(headers).build();
outputChannel.send(message);
}
catch (Throwable ex) {
logger.error("Terminating STOMP session due to failure to send message", ex);
sendErrorMessage(session, ex);
}
}
@@ -281,11 +306,14 @@ public class StompSubProtocolHandler implements SubProtocolHandler {
@Override
public void afterSessionStarted(WebSocketSession session, MessageChannel outputChannel) {
this.decoders.put(session.getId(), new BufferingStompDecoder(getMessageBufferSizeLimit()));
}
@Override
public void afterSessionEnded(WebSocketSession session, CloseStatus closeStatus, MessageChannel outputChannel) {
this.decoders.remove(session.getId());
Principal principal = session.getPrincipal();
if ((this.userSessionRegistry != null) && (principal != null)) {
String userName = resolveNameForUserSessionRegistry(principal);

View File

@@ -41,6 +41,7 @@ import org.springframework.stereotype.Controller;
import org.springframework.web.servlet.HandlerMapping;
import org.springframework.web.servlet.handler.SimpleUrlHandlerMapping;
import org.springframework.web.socket.TextMessage;
import org.springframework.web.socket.WebSocketSession;
import org.springframework.web.socket.handler.TestWebSocketSession;
import org.springframework.web.socket.messaging.StompTextMessageBuilder;
import org.springframework.web.socket.messaging.SubProtocolWebSocketHandler;
@@ -81,8 +82,11 @@ public class WebSocketMessageBrokerConfigurationSupportTests {
TestChannel channel = this.config.getBean("clientInboundChannel", TestChannel.class);
SubProtocolWebSocketHandler webSocketHandler = this.config.getBean(SubProtocolWebSocketHandler.class);
WebSocketSession session = new TestWebSocketSession("s1");
webSocketHandler.afterConnectionEstablished(session);
TextMessage textMessage = StompTextMessageBuilder.create(StompCommand.SEND).headers("destination:/foo").build();
webSocketHandler.handleMessage(new TestWebSocketSession(), textMessage);
webSocketHandler.handleMessage(session, textMessage);
Message<?> message = channel.messages.get(0);
StompHeaderAccessor headers = StompHeaderAccessor.wrap(message);

View File

@@ -20,6 +20,7 @@ import java.nio.ByteBuffer;
import java.util.Arrays;
import java.util.Collections;
import java.util.HashSet;
import java.util.List;
import org.junit.Before;
import org.junit.Test;
@@ -145,8 +146,8 @@ public class StompSubProtocolHandlerTests {
assertEquals(1, this.session.getSentMessages().size());
TextMessage textMessage = (TextMessage) this.session.getSentMessages().get(0);
Message<?> message = new StompDecoder().decode(ByteBuffer.wrap(textMessage.getPayload().getBytes()));
StompHeaderAccessor replyHeaders = StompHeaderAccessor.wrap(message);
List<Message<byte[]>> message = new StompDecoder().decode(ByteBuffer.wrap(textMessage.getPayload().getBytes()));
StompHeaderAccessor replyHeaders = StompHeaderAccessor.wrap(message.get(0));
assertEquals(StompCommand.CONNECTED, replyHeaders.getCommand());
assertEquals("1.1", replyHeaders.getVersion());
@@ -176,6 +177,7 @@ public class StompSubProtocolHandlerTests {
TextMessage textMessage = StompTextMessageBuilder.create(StompCommand.CONNECT).headers(
"login:guest", "passcode:guest", "accept-version:1.1,1.0", "heart-beat:10000,10000").build();
this.protocolHandler.afterSessionStarted(this.session, this.channel);
this.protocolHandler.handleMessageFromClient(this.session, textMessage, this.channel);
verify(this.channel).send(this.messageCaptor.capture());