From 49f9b40fba090362e44a630d828fcedf80e452e6 Mon Sep 17 00:00:00 2001 From: Jared Wiltshire Date: Mon, 3 Feb 2025 16:45:47 -0700 Subject: [PATCH 1/5] Support RFC 8441 upgrades over HTTP/2 CONNECT See gh-34362 Signed-off-by: Jared Wiltshire --- .../support/HandshakeWebSocketService.java | 24 +++++++------- .../web/socket/WebSocketHttpHeaders.java | 2 +- .../support/AbstractHandshakeHandler.java | 32 +++++++++++-------- 3 files changed, 32 insertions(+), 26 deletions(-) diff --git a/spring-webflux/src/main/java/org/springframework/web/reactive/socket/server/support/HandshakeWebSocketService.java b/spring-webflux/src/main/java/org/springframework/web/reactive/socket/server/support/HandshakeWebSocketService.java index c54f38d9bc..d7577a0251 100644 --- a/spring-webflux/src/main/java/org/springframework/web/reactive/socket/server/support/HandshakeWebSocketService.java +++ b/spring-webflux/src/main/java/org/springframework/web/reactive/socket/server/support/HandshakeWebSocketService.java @@ -205,23 +205,25 @@ public class HandshakeWebSocketService implements WebSocketService, Lifecycle { HttpMethod method = request.getMethod(); HttpHeaders headers = request.getHeaders(); - if (HttpMethod.GET != method && CONNECT_METHOD != method) { + if (HttpMethod.GET != method && !CONNECT_METHOD.equals(method)) { return Mono.error(new MethodNotAllowedException( request.getMethod(), Set.of(HttpMethod.GET, CONNECT_METHOD))); } - if (!"WebSocket".equalsIgnoreCase(headers.getUpgrade())) { - return handleBadRequest(exchange, "Invalid 'Upgrade' header: " + headers); - } + if (HttpMethod.GET == method) { + if (!"WebSocket".equalsIgnoreCase(headers.getUpgrade())) { + return handleBadRequest(exchange, "Invalid 'Upgrade' header: " + headers); + } - List connectionValue = headers.getConnection(); - if (!connectionValue.contains("Upgrade") && !connectionValue.contains("upgrade")) { - return handleBadRequest(exchange, "Invalid 'Connection' header: " + headers); - } + List connectionValue = headers.getConnection(); + if (!connectionValue.contains("Upgrade") && !connectionValue.contains("upgrade")) { + return handleBadRequest(exchange, "Invalid 'Connection' header: " + headers); + } - String key = headers.getFirst(SEC_WEBSOCKET_KEY); - if (key == null) { - return handleBadRequest(exchange, "Missing \"Sec-WebSocket-Key\" header"); + String key = headers.getFirst(SEC_WEBSOCKET_KEY); + if (key == null) { + return handleBadRequest(exchange, "Missing \"Sec-WebSocket-Key\" header"); + } } String protocol = selectProtocol(headers, handler); diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java index fa4c9037b8..1a4fac7f88 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java @@ -151,7 +151,7 @@ public class WebSocketHttpHeaders extends HttpHeaders { } /** - * Returns the value of the {@code Sec-WebSocket-Key} header. + * Returns the value of the {@code Sec-WebSocket-Protocol} header. * @return the value of the header */ public List getSecWebSocketProtocol() { diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java index acde43c3cc..fce20644c1 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java @@ -215,7 +215,7 @@ public abstract class AbstractHandshakeHandler implements HandshakeHandler, Life } try { HttpMethod httpMethod = request.getMethod(); - if (HttpMethod.GET != httpMethod && CONNECT_METHOD != httpMethod) { + if (HttpMethod.GET != httpMethod && !CONNECT_METHOD.equals(httpMethod)) { response.setStatusCode(HttpStatus.METHOD_NOT_ALLOWED); response.getHeaders().setAllow(Set.of(HttpMethod.GET, CONNECT_METHOD)); if (logger.isErrorEnabled()) { @@ -223,13 +223,15 @@ public abstract class AbstractHandshakeHandler implements HandshakeHandler, Life } return false; } - if (!"WebSocket".equalsIgnoreCase(headers.getUpgrade())) { - handleInvalidUpgradeHeader(request, response); - return false; - } - if (!headers.getConnection().contains("Upgrade") && !headers.getConnection().contains("upgrade")) { - handleInvalidConnectHeader(request, response); - return false; + if (HttpMethod.GET == httpMethod) { + if (!"WebSocket".equalsIgnoreCase(headers.getUpgrade())) { + handleInvalidUpgradeHeader(request, response); + return false; + } + if (!headers.getConnection().contains("Upgrade") && !headers.getConnection().contains("upgrade")) { + handleInvalidConnectHeader(request, response); + return false; + } } if (!isWebSocketVersionSupported(headers)) { handleWebSocketVersionNotSupported(request, response); @@ -239,13 +241,15 @@ public abstract class AbstractHandshakeHandler implements HandshakeHandler, Life response.setStatusCode(HttpStatus.FORBIDDEN); return false; } - String wsKey = headers.getSecWebSocketKey(); - if (wsKey == null) { - if (logger.isErrorEnabled()) { - logger.error("Missing \"Sec-WebSocket-Key\" header"); + if (HttpMethod.GET == httpMethod) { + String wsKey = headers.getSecWebSocketKey(); + if (wsKey == null) { + if (logger.isErrorEnabled()) { + logger.error("Missing \"Sec-WebSocket-Key\" header"); + } + response.setStatusCode(HttpStatus.BAD_REQUEST); + return false; } - response.setStatusCode(HttpStatus.BAD_REQUEST); - return false; } } catch (IOException ex) { From ceffda7874ddc6832c649554f56ce6769049543f Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Mon, 10 Feb 2025 09:44:09 +0000 Subject: [PATCH 2/5] Polishing contribution Closes gh-34362 --- .../web/socket/WebSocketHttpHeaders.java | 2 +- .../support/AbstractHandshakeHandler.java | 21 +++++++++---------- 2 files changed, 11 insertions(+), 12 deletions(-) diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java index 1a4fac7f88..eed9f581e8 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2023 the original author or authors. + * Copyright 2002-2025 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. diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java index fce20644c1..a233c83a2c 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/server/support/AbstractHandshakeHandler.java @@ -228,10 +228,19 @@ public abstract class AbstractHandshakeHandler implements HandshakeHandler, Life handleInvalidUpgradeHeader(request, response); return false; } - if (!headers.getConnection().contains("Upgrade") && !headers.getConnection().contains("upgrade")) { + List connectionValue = headers.getConnection(); + if (!connectionValue.contains("Upgrade") && !connectionValue.contains("upgrade")) { handleInvalidConnectHeader(request, response); return false; } + String key = headers.getSecWebSocketKey(); + if (key == null) { + if (logger.isErrorEnabled()) { + logger.error("Missing \"Sec-WebSocket-Key\" header"); + } + response.setStatusCode(HttpStatus.BAD_REQUEST); + return false; + } } if (!isWebSocketVersionSupported(headers)) { handleWebSocketVersionNotSupported(request, response); @@ -241,16 +250,6 @@ public abstract class AbstractHandshakeHandler implements HandshakeHandler, Life response.setStatusCode(HttpStatus.FORBIDDEN); return false; } - if (HttpMethod.GET == httpMethod) { - String wsKey = headers.getSecWebSocketKey(); - if (wsKey == null) { - if (logger.isErrorEnabled()) { - logger.error("Missing \"Sec-WebSocket-Key\" header"); - } - response.setStatusCode(HttpStatus.BAD_REQUEST); - return false; - } - } } catch (IOException ex) { throw new HandshakeFailureException( From c41b0140cdabd8fbf7bff43a7b8327054fc9123b Mon Sep 17 00:00:00 2001 From: Branden Clark Date: Tue, 28 Jan 2025 17:38:45 -0800 Subject: [PATCH 3/5] Check hasNext on sessionIds in UserDestinationResult See gh-34333 Signed-off-by: Branden Clark --- .../user/UserDestinationMessageHandler.java | 6 ++---- .../UserDestinationMessageHandlerTests.java | 21 +++++++++++++++++++ 2 files changed, 23 insertions(+), 4 deletions(-) diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java index 9c8d24faab..4693f6c211 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java @@ -20,7 +20,6 @@ import java.util.Arrays; import java.util.Iterator; import java.util.List; import java.util.Map; -import java.util.Set; import java.util.concurrent.ConcurrentHashMap; import org.apache.commons.logging.Log; @@ -282,11 +281,10 @@ public class UserDestinationMessageHandler implements MessageHandler, SmartLifec } public void send(UserDestinationResult destinationResult, Message message) throws MessagingException { - Set sessionIds = destinationResult.getSessionIds(); - Iterator itr = (sessionIds != null ? sessionIds.iterator() : null); + Iterator itr = destinationResult.getSessionIds().iterator(); for (String target : destinationResult.getTargetDestinations()) { - String sessionId = (itr != null ? itr.next() : null); + String sessionId = (itr != null && itr.hasNext() ? itr.next() : null); getTemplateToUse(sessionId).send(target, message); } } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java index e1ec61dd89..c22d1f45e3 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java @@ -17,6 +17,7 @@ package org.springframework.messaging.simp.user; import java.nio.charset.StandardCharsets; +import java.util.Set; import org.junit.jupiter.api.Test; import org.mockito.ArgumentCaptor; @@ -98,6 +99,26 @@ class UserDestinationMessageHandlerTests { assertThat(accessor.getFirstNativeHeader(ORIGINAL_DESTINATION)).isEqualTo("/user/queue/foo"); } + @Test + @SuppressWarnings("rawtypes") + void handleMessageWithoutSessionIds() { + UserDestinationResolver resolver = mock(); + Message message = createWith(SimpMessageType.MESSAGE, "joe", null, "/user/joe/queue/foo"); + UserDestinationResult result = new UserDestinationResult("/queue/foo-user123", Set.of("/queue/foo-user123"), "/user/queue/foo", "joe"); + given(resolver.resolveDestination(message)).willReturn(result); + + given(this.brokerChannel.send(Mockito.any(Message.class))).willReturn(true); + UserDestinationMessageHandler handler = new UserDestinationMessageHandler(new StubMessageChannel(), this.brokerChannel, resolver); + handler.handleMessage(message); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Message.class); + Mockito.verify(this.brokerChannel).send(captor.capture()); + + SimpMessageHeaderAccessor accessor = SimpMessageHeaderAccessor.wrap(captor.getValue()); + assertThat(accessor.getDestination()).isEqualTo("/queue/foo-user123"); + assertThat(accessor.getFirstNativeHeader(ORIGINAL_DESTINATION)).isEqualTo("/user/queue/foo"); + } + @Test @SuppressWarnings("rawtypes") void handleMessageWithoutActiveSession() { From ccdaed594eaa44b9d3ba8a0df7d9c13ffcee443e Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Mon, 10 Feb 2025 10:09:21 +0000 Subject: [PATCH 4/5] Polishing contribution Closes gh-34333 --- .../simp/user/UserDestinationMessageHandler.java | 11 +++++------ .../messaging/simp/user/UserDestinationResult.java | 9 ++++++--- .../simp/user/UserDestinationMessageHandlerTests.java | 2 +- 3 files changed, 12 insertions(+), 10 deletions(-) diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java index 4693f6c211..791602c0c1 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationMessageHandler.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2024 the original author or authors. + * Copyright 2002-2025 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. @@ -280,11 +280,10 @@ public class UserDestinationMessageHandler implements MessageHandler, SmartLifec return this.messagingTemplate; } - public void send(UserDestinationResult destinationResult, Message message) throws MessagingException { - Iterator itr = destinationResult.getSessionIds().iterator(); - - for (String target : destinationResult.getTargetDestinations()) { - String sessionId = (itr != null && itr.hasNext() ? itr.next() : null); + public void send(UserDestinationResult result, Message message) throws MessagingException { + Iterator itr = result.getSessionIds().iterator(); + for (String target : result.getTargetDestinations()) { + String sessionId = (itr.hasNext() ? itr.next() : null); getTemplateToUse(sessionId).send(target, message); } } diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationResult.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationResult.java index aa74978e1f..2abcbe93e0 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationResult.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/user/UserDestinationResult.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2023 the original author or authors. + * Copyright 2002-2025 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. @@ -44,7 +44,11 @@ public class UserDestinationResult { private final Set sessionIds; - public UserDestinationResult(String sourceDestination, Set targetDestinations, + /** + * Main constructor. + */ + public UserDestinationResult( + String sourceDestination, Set targetDestinations, String subscribeDestination, @Nullable String user) { this(sourceDestination, targetDestinations, subscribeDestination, user, null); @@ -114,7 +118,6 @@ public class UserDestinationResult { /** * Return the session id for the targetDestination. */ - @Nullable public Set getSessionIds() { return this.sessionIds; } diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java index c22d1f45e3..cfcfc1d08d 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/user/UserDestinationMessageHandlerTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2024 the original author or authors. + * Copyright 2002-2025 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. From 7a0fe7d14faccc4574c3ee49ae92d657aacaa2fe Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Mon, 10 Feb 2025 11:14:02 +0000 Subject: [PATCH 5/5] WebAsyncManager wraps disconnected client errors If the Servlet container delegates a disconnected client error via AsyncListener#onError, wrap it as AsyncRequestNotUsableException for more targeted and consistent handling of such errors. Closes gh-34363 --- .../request/async/WebAsyncManager.java | 11 +++++- .../async/WebAsyncManagerErrorTests.java | 34 ++++++++++++++++++- 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/spring-web/src/main/java/org/springframework/web/context/request/async/WebAsyncManager.java b/spring-web/src/main/java/org/springframework/web/context/request/async/WebAsyncManager.java index 2cb1ae18ea..00479d483e 100644 --- a/spring-web/src/main/java/org/springframework/web/context/request/async/WebAsyncManager.java +++ b/spring-web/src/main/java/org/springframework/web/context/request/async/WebAsyncManager.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2024 the original author or authors. + * Copyright 2002-2025 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. @@ -34,6 +34,7 @@ import org.springframework.lang.Nullable; import org.springframework.util.Assert; import org.springframework.web.context.request.RequestAttributes; import org.springframework.web.context.request.async.DeferredResult.DeferredResultHandler; +import org.springframework.web.util.DisconnectedClientHelper; /** * The central class for managing asynchronous request processing, mainly intended @@ -350,6 +351,10 @@ public final class WebAsyncManager { if (logger.isDebugEnabled()) { logger.debug("Servlet container error notification for " + formatUri(this.asyncWebRequest) + ": " + ex); } + if (DisconnectedClientHelper.isClientDisconnectedException(ex)) { + ex = new AsyncRequestNotUsableException( + "Servlet container error notification for disconnected client", ex); + } Object result = interceptorChain.triggerAfterError(this.asyncWebRequest, callable, ex); result = (result != CallableProcessingInterceptor.RESULT_NONE ? result : ex); setConcurrentResultAndDispatch(result); @@ -442,6 +447,10 @@ public final class WebAsyncManager { if (logger.isDebugEnabled()) { logger.debug("Servlet container error notification for " + formatUri(this.asyncWebRequest)); } + if (DisconnectedClientHelper.isClientDisconnectedException(ex)) { + ex = new AsyncRequestNotUsableException( + "Servlet container error notification for disconnected client", ex); + } try { interceptorChain.triggerAfterError(this.asyncWebRequest, deferredResult, ex); synchronized (WebAsyncManager.this) { diff --git a/spring-web/src/test/java/org/springframework/web/context/request/async/WebAsyncManagerErrorTests.java b/spring-web/src/test/java/org/springframework/web/context/request/async/WebAsyncManagerErrorTests.java index 542638dcbf..d11d6714c0 100644 --- a/spring-web/src/test/java/org/springframework/web/context/request/async/WebAsyncManagerErrorTests.java +++ b/spring-web/src/test/java/org/springframework/web/context/request/async/WebAsyncManagerErrorTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2024 the original author or authors. + * Copyright 2002-2025 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. @@ -16,6 +16,7 @@ package org.springframework.web.context.request.async; +import java.io.IOException; import java.util.concurrent.Callable; import jakarta.servlet.AsyncEvent; @@ -152,6 +153,22 @@ class WebAsyncManagerErrorTests { verify(interceptor).beforeConcurrentHandling(this.asyncWebRequest, callable); } + @Test // gh-34363 + void startCallableProcessingDisconnectedClient() throws Exception { + StubCallable callable = new StubCallable(); + this.asyncManager.startCallableProcessing(callable); + + IOException ex = new IOException("broken pipe"); + AsyncEvent event = new AsyncEvent(new MockAsyncContext(this.servletRequest, this.servletResponse), ex); + this.asyncWebRequest.onError(event); + + MockAsyncContext asyncContext = (MockAsyncContext) this.servletRequest.getAsyncContext(); + assertThat(this.asyncManager.hasConcurrentResult()).isTrue(); + assertThat(this.asyncManager.getConcurrentResult()) + .as("Disconnected client error not wrapped AsyncRequestNotUsableException") + .isOfAnyClassIn(AsyncRequestNotUsableException.class); + } + @Test void startDeferredResultProcessingErrorAndComplete() throws Exception { @@ -259,6 +276,21 @@ class WebAsyncManagerErrorTests { assertThat(((MockAsyncContext) this.servletRequest.getAsyncContext()).getDispatchedPath()).isEqualTo("/test"); } + @Test // gh-34363 + void startDeferredResultProcessingDisconnectedClient() throws Exception { + DeferredResult deferredResult = new DeferredResult<>(); + this.asyncManager.startDeferredResultProcessing(deferredResult); + + IOException ex = new IOException("broken pipe"); + AsyncEvent event = new AsyncEvent(new MockAsyncContext(this.servletRequest, this.servletResponse), ex); + this.asyncWebRequest.onError(event); + + assertThat(this.asyncManager.hasConcurrentResult()).isTrue(); + assertThat(deferredResult.getResult()) + .as("Disconnected client error not wrapped AsyncRequestNotUsableException") + .isOfAnyClassIn(AsyncRequestNotUsableException.class); + } + private static final class StubCallable implements Callable { @Override