From 26309838ba3727a7358c331325df4f3b325c4674 Mon Sep 17 00:00:00 2001 From: Brian Clozel Date: Fri, 21 Mar 2014 10:00:11 +0100 Subject: [PATCH] Set custom handshakeHandler for XML sockjs config Prior to this commit, configuring a custom handshakeHandler when setting up a stomp-endpoint with SockJS would not be taken into account: This commit fixes this by creating and registering a WebsocketTransportHandler (with this handshakeHandler) as a transportHandler override for the SockJSService. Issue: SPR-11568 --- .../MessageBrokerBeanDefinitionParser.java | 5 ++--- .../config/WebSocketNamespaceUtils.java | 12 +++++++++++- ...essageBrokerBeanDefinitionParserTests.java | 19 ++++++++++++++++++- .../config/websocket-config-broker-relay.xml | 1 - 4 files changed, 31 insertions(+), 6 deletions(-) diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParser.java b/spring-websocket/src/main/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParser.java index c2cf38e322..68783e6ad3 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParser.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParser.java @@ -249,9 +249,6 @@ class MessageBrokerBeanDefinitionParser implements BeanDefinitionParser { RootBeanDefinition httpRequestHandlerDef; - RuntimeBeanReference handshakeHandler = - WebSocketNamespaceUtils.registerHandshakeHandler(stompEndpointElement, parserCxt, source); - RuntimeBeanReference sockJsService = WebSocketNamespaceUtils.registerSockJsService( stompEndpointElement, SOCKJS_SCHEDULER_BEAN_NAME, parserCxt, source); @@ -262,6 +259,8 @@ class MessageBrokerBeanDefinitionParser implements BeanDefinitionParser { httpRequestHandlerDef = new RootBeanDefinition(SockJsHttpRequestHandler.class, cavs, null); } else { + RuntimeBeanReference handshakeHandler = + WebSocketNamespaceUtils.registerHandshakeHandler(stompEndpointElement, parserCxt, source); ConstructorArgumentValues cavs = new ConstructorArgumentValues(); cavs.addIndexedArgumentValue(0, subProtocolWebSocketHandler); if(handshakeHandler != null) { diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/config/WebSocketNamespaceUtils.java b/spring-websocket/src/main/java/org/springframework/web/socket/config/WebSocketNamespaceUtils.java index 420a810bb4..aa99dcce84 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/config/WebSocketNamespaceUtils.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/config/WebSocketNamespaceUtils.java @@ -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. @@ -29,6 +29,7 @@ import org.springframework.util.xml.DomUtils; import org.springframework.web.socket.server.support.DefaultHandshakeHandler; import org.springframework.web.socket.sockjs.transport.TransportHandlingSockJsService; import org.springframework.web.socket.sockjs.transport.handler.DefaultSockJsService; +import org.springframework.web.socket.sockjs.transport.handler.WebSocketTransportHandler; /** * Provides utility methods for parsing common WebSocket XML namespace elements. @@ -61,6 +62,8 @@ class WebSocketNamespaceUtils { Element sockJsElement = DomUtils.getChildElementByTagName(element, "sockjs"); if (sockJsElement != null) { + Element handshakeHandlerElement = DomUtils.getChildElementByTagName(element, "handshake-handler"); + RootBeanDefinition sockJsServiceDef = new RootBeanDefinition(DefaultSockJsService.class); sockJsServiceDef.setSource(source); @@ -82,6 +85,13 @@ class WebSocketNamespaceUtils { } ManagedList transportHandlersList = parseBeanSubElements(transportHandlersElement, parserContext); sockJsServiceDef.getConstructorArgumentValues().addIndexedArgumentValue(1, transportHandlersList); + } else if(handshakeHandlerElement != null){ + RuntimeBeanReference handshakeHandlerRef = new RuntimeBeanReference(handshakeHandlerElement.getAttribute("ref")); + + RootBeanDefinition wsTransportHandler = new RootBeanDefinition(WebSocketTransportHandler.class); + wsTransportHandler.setSource(source); + wsTransportHandler.getConstructorArgumentValues().addIndexedArgumentValue(0, handshakeHandlerRef); + sockJsServiceDef.getConstructorArgumentValues().addIndexedArgumentValue(1, wsTransportHandler); } String attrValue = sockJsElement.getAttribute("name"); diff --git a/spring-websocket/src/test/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParserTests.java b/spring-websocket/src/test/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParserTests.java index 2efd503057..8e958c27f8 100644 --- a/spring-websocket/src/test/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParserTests.java +++ b/spring-websocket/src/test/java/org/springframework/web/socket/config/MessageBrokerBeanDefinitionParserTests.java @@ -18,13 +18,13 @@ package org.springframework.web.socket.config; import java.util.ArrayList; import java.util.Arrays; -import java.util.Collections; import java.util.List; import org.hamcrest.Matchers; import org.junit.Before; import org.junit.Test; +import org.springframework.beans.DirectFieldAccessor; import org.springframework.beans.factory.NoSuchBeanDefinitionException; import org.springframework.beans.factory.xml.XmlBeanDefinitionReader; import org.springframework.core.io.ClassPathResource; @@ -55,8 +55,12 @@ import org.springframework.web.socket.WebSocketHandler; import org.springframework.web.socket.handler.WebSocketHandlerDecorator; import org.springframework.web.socket.messaging.StompSubProtocolHandler; import org.springframework.web.socket.messaging.SubProtocolWebSocketHandler; +import org.springframework.web.socket.server.HandshakeHandler; import org.springframework.web.socket.server.support.WebSocketHttpRequestHandler; import org.springframework.web.socket.sockjs.support.SockJsHttpRequestHandler; +import org.springframework.web.socket.sockjs.transport.TransportType; +import org.springframework.web.socket.sockjs.transport.handler.DefaultSockJsService; +import org.springframework.web.socket.sockjs.transport.handler.WebSocketTransportHandler; import static org.junit.Assert.*; import static org.junit.Assert.assertEquals; @@ -93,6 +97,12 @@ public class MessageBrokerBeanDefinitionParserTests { assertThat(httpRequestHandler, Matchers.instanceOf(WebSocketHttpRequestHandler.class)); WebSocketHttpRequestHandler wsHttpRequestHandler = (WebSocketHttpRequestHandler) httpRequestHandler; + + HandshakeHandler handshakeHandler = (HandshakeHandler) + new DirectFieldAccessor(wsHttpRequestHandler).getPropertyValue("handshakeHandler"); + assertNotNull(handshakeHandler); + assertTrue(handshakeHandler instanceof TestHandshakeHandler); + WebSocketHandler wsHandler = unwrapWebSocketHandler(wsHttpRequestHandler.getWebSocketHandler()); assertNotNull(wsHandler); assertThat(wsHandler, Matchers.instanceOf(SubProtocolWebSocketHandler.class)); @@ -113,6 +123,13 @@ public class MessageBrokerBeanDefinitionParserTests { assertNotNull(wsHandler); assertThat(wsHandler, Matchers.instanceOf(SubProtocolWebSocketHandler.class)); assertNotNull(sockJsHttpRequestHandler.getSockJsService()); + assertThat(sockJsHttpRequestHandler.getSockJsService(), Matchers.instanceOf(DefaultSockJsService.class)); + + DefaultSockJsService defaultSockJsService = (DefaultSockJsService) sockJsHttpRequestHandler.getSockJsService(); + WebSocketTransportHandler wsTransportHandler = (WebSocketTransportHandler) defaultSockJsService + .getTransportHandlers().get(TransportType.WEBSOCKET); + assertNotNull(wsTransportHandler.getHandshakeHandler()); + assertThat(wsTransportHandler.getHandshakeHandler(), Matchers.instanceOf(TestHandshakeHandler.class)); UserSessionRegistry userSessionRegistry = this.appContext.getBean(UserSessionRegistry.class); assertNotNull(userSessionRegistry); diff --git a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-relay.xml b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-relay.xml index a5fd205b84..4a6f0fa85c 100644 --- a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-relay.xml +++ b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-relay.xml @@ -6,7 +6,6 @@ -