diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/server/DefaultHandshakeHandler.java b/spring-websocket/src/main/java/org/springframework/web/socket/server/DefaultHandshakeHandler.java index 660dac28ea..80dcc5aa85 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/server/DefaultHandshakeHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/server/DefaultHandshakeHandler.java @@ -17,16 +17,12 @@ package org.springframework.web.socket.server; import java.io.IOException; -import java.nio.charset.Charset; -import java.security.MessageDigest; -import java.security.NoSuchAlgorithmException; import java.util.ArrayList; +import java.util.Arrays; import java.util.Collections; import java.util.List; import java.util.Map; -import javax.xml.bind.DatatypeConverter; - import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.springframework.beans.BeanUtils; @@ -39,24 +35,38 @@ import org.springframework.util.StringUtils; import org.springframework.web.socket.WebSocketHandler; /** - * A default implemnetation of {@link HandshakeHandler}. + * A default {@link HandshakeHandler} implementation. Performs initial validation of the + * WebSocket handshake request -- possibly rejecting it through the appropriate HTTP + * status code -- while also allowing sub-classes to override various parts of the + * negotiation process (e.g. origin validation, sub-protocol negotiation, etc). * - *

A container-specific {@link RequestUpgradeStrategy} is required since standard Java - * WebSocket currently does not provide a way to initiate a WebSocket handshake. - * Currently available are implementations for Tomcat and GlassFish. + *

+ * If the negotiation succeeds, the actual upgrade is delegated to a server-specific + * {@link RequestUpgradeStrategy}, which will update the response as necessary and + * initialize the WebSocket. Currently supported servers are Tomcat 7 and 8, Jetty 9, and + * Glassfish 4. * * @author Rossen Stoyanchev * @since 4.0 */ public class DefaultHandshakeHandler implements HandshakeHandler { - private static final String GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; - protected Log logger = LogFactory.getLog(getClass()); + private static final boolean tomcatWsPresent = ClassUtils.isPresent( + "org.apache.tomcat.websocket.server.WsHttpUpgradeHandler", HandshakeHandler.class.getClassLoader()); + + private static final boolean jettyWsPresent = ClassUtils.isPresent( + "org.eclipse.jetty.websocket.server.WebSocketServerFactory", HandshakeHandler.class.getClassLoader()); + + private static final boolean glassFishWsPresent = ClassUtils.isPresent( + "org.glassfish.tyrus.servlet.TyrusHttpUpgradeHandler", HandshakeHandler.class.getClassLoader()); + + + private final RequestUpgradeStrategy requestUpgradeStrategy; + private final List supportedProtocols = new ArrayList(); - private final RequestUpgradeStrategy requestUpgradeStrategy; /** @@ -66,7 +76,30 @@ public class DefaultHandshakeHandler implements HandshakeHandler { * @throws IllegalStateException if no {@link RequestUpgradeStrategy} can be found. */ public DefaultHandshakeHandler() { - this.requestUpgradeStrategy = new RequestUpgradeStrategyFactory().create(); + this(initRequestUpgradeStrategy()); + } + + private static RequestUpgradeStrategy initRequestUpgradeStrategy() { + String className; + if (tomcatWsPresent) { + className = "org.springframework.web.socket.server.support.TomcatRequestUpgradeStrategy"; + } + else if (jettyWsPresent) { + className = "org.springframework.web.socket.server.support.JettyRequestUpgradeStrategy"; + } + else if (glassFishWsPresent) { + className = "org.springframework.web.socket.server.support.GlassFishRequestUpgradeStrategy"; + } + else { + throw new IllegalStateException("No suitable " + RequestUpgradeStrategy.class.getSimpleName()); + } + try { + Class clazz = ClassUtils.forName(className, DefaultHandshakeHandler.class.getClassLoader()); + return (RequestUpgradeStrategy) BeanUtils.instantiateClass(clazz.getConstructor()); + } + catch (Throwable t) { + throw new IllegalStateException("Failed to instantiate " + className, t); + } } /** @@ -101,7 +134,9 @@ public class DefaultHandshakeHandler implements HandshakeHandler { public final boolean doHandshake(ServerHttpRequest request, ServerHttpResponse response, WebSocketHandler webSocketHandler, Map attributes) throws IOException, HandshakeFailureException { - logger.debug("Starting handshake for " + request.getURI()); + if (logger.isDebugEnabled()) { + logger.debug("Initiating handshake for " + request.getURI() + ", headers=" + request.getHeaders()); + } if (!HttpMethod.GET.equals(request.getMethod())) { response.setStatusCode(HttpStatus.METHOD_NOT_ALLOWED); @@ -136,19 +171,8 @@ public class DefaultHandshakeHandler implements HandshakeHandler { String selectedProtocol = selectProtocol(request.getHeaders().getSecWebSocketProtocol()); // TODO: select extensions - logger.debug("Upgrading HTTP request"); - - response.setStatusCode(HttpStatus.SWITCHING_PROTOCOLS); - response.getHeaders().setUpgrade("WebSocket"); - response.getHeaders().setConnection("Upgrade"); - response.getHeaders().setSecWebSocketProtocol(selectedProtocol); - response.getHeaders().setSecWebSocketAccept(getWebSocketKeyHash(wsKey)); - // TODO: response.getHeaders().setSecWebSocketExtensions(extensions); - - response.flush(); - - if (logger.isTraceEnabled()) { - logger.trace("Upgrading with " + webSocketHandler); + if (logger.isDebugEnabled()) { + logger.debug("Upgrading request"); } this.requestUpgradeStrategy.upgrade(request, response, selectedProtocol, webSocketHandler, attributes); @@ -169,12 +193,19 @@ public class DefaultHandshakeHandler implements HandshakeHandler { } protected boolean isWebSocketVersionSupported(ServerHttpRequest request) { - String requestedVersion = request.getHeaders().getSecWebSocketVersion(); - for (String supportedVersion : getSupportedVerions()) { - if (supportedVersion.equals(requestedVersion)) { + String version = request.getHeaders().getSecWebSocketVersion(); + String[] supportedVersions = getSupportedVerions(); + if (logger.isDebugEnabled()) { + logger.debug("Requested version=" + version + ", supported=" + Arrays.toString(supportedVersions)); + } + for (String supportedVersion : supportedVersions) { + if (supportedVersion.trim().equals(version)) { return true; } } + if (logger.isDebugEnabled()) { + logger.debug("Version=" + version + " is not a supported version"); + } return false; } @@ -218,51 +249,4 @@ public class DefaultHandshakeHandler implements HandshakeHandler { return null; } - private String getWebSocketKeyHash(String key) throws HandshakeFailureException { - try { - MessageDigest digest = MessageDigest.getInstance("SHA1"); - byte[] bytes = digest.digest((key + GUID).getBytes(Charset.forName("ISO-8859-1"))); - return DatatypeConverter.printBase64Binary(bytes); - } - catch (NoSuchAlgorithmException ex) { - throw new HandshakeFailureException("Failed to generate value for Sec-WebSocket-Key header", ex); - } - } - - - private static class RequestUpgradeStrategyFactory { - - private static final boolean tomcatWebSocketPresent = ClassUtils.isPresent( - "org.apache.tomcat.websocket.server.WsHttpUpgradeHandler", DefaultHandshakeHandler.class.getClassLoader()); - - private static final boolean glassFishWebSocketPresent = ClassUtils.isPresent( - "org.glassfish.tyrus.servlet.TyrusHttpUpgradeHandler", DefaultHandshakeHandler.class.getClassLoader()); - - private static final boolean jettyWebSocketPresent = ClassUtils.isPresent( - "org.eclipse.jetty.websocket.server.UpgradeContext", DefaultHandshakeHandler.class.getClassLoader()); - - private RequestUpgradeStrategy create() { - String className; - if (tomcatWebSocketPresent) { - className = "org.springframework.web.socket.server.support.TomcatRequestUpgradeStrategy"; - } - else if (glassFishWebSocketPresent) { - className = "org.springframework.web.socket.server.support.GlassFishRequestUpgradeStrategy"; - } - else if (jettyWebSocketPresent) { - className = "org.springframework.web.socket.server.support.JettyRequestUpgradeStrategy"; - } - else { - throw new IllegalStateException("No suitable " + RequestUpgradeStrategy.class.getSimpleName()); - } - try { - Class clazz = ClassUtils.forName(className, DefaultHandshakeHandler.class.getClassLoader()); - return (RequestUpgradeStrategy) BeanUtils.instantiateClass(clazz.getConstructor()); - } - catch (Throwable t) { - throw new IllegalStateException("Failed to instantiate " + className, t); - } - } - } - } diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/server/RequestUpgradeStrategy.java b/spring-websocket/src/main/java/org/springframework/web/socket/server/RequestUpgradeStrategy.java index 3dc5d83aed..9a8e360d35 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/server/RequestUpgradeStrategy.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/server/RequestUpgradeStrategy.java @@ -24,14 +24,14 @@ import org.springframework.http.server.ServerHttpResponse; import org.springframework.web.socket.WebSocketHandler; /** - * A strategy for performing container-specific steps to upgrade an HTTP request during a - * WebSocket handshake. Intended for use within {@link HandshakeHandler} implementations. + * A server-specific strategy for performing the actual upgrade to a WebSocket exchange. * * @author Rossen Stoyanchev * @since 4.0 */ public interface RequestUpgradeStrategy { + /** * Return the supported WebSocket protocol versions. */ diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/GlassFishRequestUpgradeStrategy.java b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/GlassFishRequestUpgradeStrategy.java index be8048739e..cef4dc7ae4 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/GlassFishRequestUpgradeStrategy.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/GlassFishRequestUpgradeStrategy.java @@ -25,7 +25,6 @@ import java.util.Random; import javax.servlet.ServletException; import javax.servlet.http.HttpServletRequest; import javax.servlet.http.HttpServletResponse; -import javax.servlet.http.HttpServletResponseWrapper; import javax.websocket.DeploymentException; import javax.websocket.Endpoint; @@ -41,7 +40,6 @@ import org.glassfish.tyrus.websockets.WebSocketApplication; import org.glassfish.tyrus.websockets.WebSocketEngine; import org.glassfish.tyrus.websockets.WebSocketEngine.WebSocketHolderListener; import org.springframework.http.HttpHeaders; -import org.springframework.http.HttpStatus; import org.springframework.http.server.ServerHttpRequest; import org.springframework.http.server.ServerHttpResponse; import org.springframework.http.server.ServletServerHttpRequest; @@ -53,9 +51,9 @@ import org.springframework.util.StringUtils; import org.springframework.web.socket.server.HandshakeFailureException; import org.springframework.web.socket.server.endpoint.ServerEndpointRegistration; + /** - * GlassFish support for upgrading an {@link HttpServletRequest} during a WebSocket - * handshake. + * GlassFish support for upgrading a request during a WebSocket handshake. * * @author Rossen Stoyanchev * @since 4.0 @@ -79,25 +77,22 @@ public class GlassFishRequestUpgradeStrategy extends AbstractStandardUpgradeStra Assert.isTrue(response instanceof ServletServerHttpResponse); HttpServletResponse servletResponse = ((ServletServerHttpResponse) response).getServletResponse(); - servletResponse = new AlreadyUpgradedResponseWrapper(servletResponse); - WebSocketApplication wsApp = createTyrusEndpoint(servletRequest, endpoint, selectedProtocol); - WebSocketEngine engine = WebSocketEngine.getEngine(); + WebSocketApplication webSocketApplication = createTyrusEndpoint(servletRequest, endpoint, selectedProtocol); + WebSocketEngine webSocketEngine = WebSocketEngine.getEngine(); try { - engine.register(wsApp); + webSocketEngine.register(webSocketApplication); } catch (DeploymentException ex) { throw new HandshakeFailureException("Failed to deploy endpoint in GlassFish", ex); } try { - if (!performUpgrade(servletRequest, servletResponse, request.getHeaders(), wsApp)) { - throw new HandshakeFailureException("Failed to upgrade HttpServletRequest"); - } + performUpgrade(servletRequest, servletResponse, request.getHeaders(), webSocketApplication); } finally { - engine.unregister(wsApp); + webSocketEngine.unregister(webSocketApplication); } } @@ -133,11 +128,8 @@ public class GlassFishRequestUpgradeStrategy extends AbstractStandardUpgradeStra private WebSocketApplication createTyrusEndpoint(HttpServletRequest request, Endpoint endpoint, String selectedProtocol) { - // Use randomized path - String requestUri = request.getRequestURI(); - String randomValue = String.valueOf(random.nextLong()); - String endpointPath = requestUri.endsWith("/") ? requestUri + randomValue : requestUri + "/" + randomValue; - + // shouldn't matter for processing but must be unique + String endpointPath = "/" + random.nextLong(); ServerEndpointRegistration endpointConfig = new ServerEndpointRegistration(endpointPath, endpoint); endpointConfig.setSubprotocols(Arrays.asList(selectedProtocol)); @@ -159,21 +151,4 @@ public class GlassFishRequestUpgradeStrategy extends AbstractStandardUpgradeStra } } - - private static class AlreadyUpgradedResponseWrapper extends HttpServletResponseWrapper { - - public AlreadyUpgradedResponseWrapper(HttpServletResponse response) { - super(response); - } - - @Override - public void setStatus(int sc) { - Assert.isTrue(sc == HttpStatus.SWITCHING_PROTOCOLS.value(), "Unexpected status code " + sc); - } - @Override - public void addHeader(String name, String value) { - // ignore - } - } - } diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSession.java b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSession.java index 1fbd5e1de3..8767f21965 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSession.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSession.java @@ -255,12 +255,13 @@ public abstract class AbstractSockJsSession implements WebSocketSession { delegateError(ex); } catch (Throwable delegateEx) { - try { - close(closeStatus); - } - catch (Throwable closeEx) { - // ignore - } + // ignore + } + try { + close(closeStatus); + } + catch (Throwable closeEx) { + // ignore } } diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/WebSocketServerSockJsSession.java b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/WebSocketServerSockJsSession.java index ac3cbfcf6a..678a6c6d4f 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/WebSocketServerSockJsSession.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/session/WebSocketServerSockJsSession.java @@ -124,7 +124,7 @@ public class WebSocketServerSockJsSession extends AbstractSockJsSession try { messages = getSockJsServiceConfig().getMessageCodec().decode(payload); } - catch (IOException ex) { + catch (Throwable ex) { logger.error("Broken data received. Terminating WebSocket connection abruptly", ex); tryCloseWithSockJsTransportError(ex, CloseStatus.BAD_DATA); return; diff --git a/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSessionTests.java b/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSessionTests.java index f4a76713bb..093318d471 100644 --- a/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSessionTests.java +++ b/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/session/AbstractSockJsSessionTests.java @@ -215,6 +215,17 @@ public class AbstractSockJsSessionTests extends BaseAbstractSockJsSessionTests