From 0044c4c8a39d6546e44966ba912aab89962c13b9 Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Thu, 25 Jun 2026 16:09:55 +0100 Subject: [PATCH] Improve remoteAddress check in TransportHandlingSockJsService Closes gh-36904 --- .../TransportHandlingSockJsService.java | 15 ++++++++- .../handler/DefaultSockJsServiceTests.java | 31 +++++++++++++++++-- 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/TransportHandlingSockJsService.java b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/TransportHandlingSockJsService.java index fc0cb7318c9..1b3625b6e2a 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/TransportHandlingSockJsService.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/sockjs/transport/TransportHandlingSockJsService.java @@ -321,7 +321,7 @@ public class TransportHandlingSockJsService extends AbstractSockJsService implem return; } InetSocketAddress remoteAddress = session.getRemoteAddress(); - if (remoteAddress != null && !remoteAddress.equals(request.getRemoteAddress())) { + if (remoteAddress != null && !isSameAddress(remoteAddress, request.getRemoteAddress())) { logger.debug("The remote address for the session and the request do not match."); response.setStatusCode(HttpStatus.NOT_FOUND); return; @@ -366,6 +366,19 @@ public class TransportHandlingSockJsService extends AbstractSockJsService implem } } + private boolean isSameAddress(InetSocketAddress address, InetSocketAddress that) { + // InetSocketAddress#equals minus port checks, which can vary by requests + if (address.getAddress() != null) { + return address.getAddress().equals(that.getAddress()); + } + else if (address.getHostName() != null) { + return (that.getAddress() == null && address.getHostName().equalsIgnoreCase(that.getHostName())); + } + else { + return (that.getAddress() == null) && (that.getHostName() == null); + } + } + @Override protected boolean validateRequest(String serverId, String sessionId, String transport) { if (!super.validateRequest(serverId, sessionId, transport)) { diff --git a/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/handler/DefaultSockJsServiceTests.java b/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/handler/DefaultSockJsServiceTests.java index 70f62314ce6..5cda438859a 100644 --- a/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/handler/DefaultSockJsServiceTests.java +++ b/spring-websocket/src/test/java/org/springframework/web/socket/sockjs/transport/handler/DefaultSockJsServiceTests.java @@ -291,8 +291,9 @@ class DefaultSockJsServiceTests extends AbstractHttpRequestTests { assertThat(this.servletResponse.getStatus()).isEqualTo(200); verify(this.xhrHandler).handleRequest(this.request, this.response, this.wsHandler, this.session); - this.session.setRemoteAddress(new InetSocketAddress("127.0.0.1:8080", 8080)); - this.servletRequest.setRemoteAddr("127.0.0.1:9090"); + this.session.setRemoteAddress(new InetSocketAddress("0.0.0.0.1", 54001)); + this.servletRequest.setRemoteAddr("0.0.0.0.2"); + this.servletRequest.setRemotePort(54001); resetResponse(); reset(this.xhrSendHandler); @@ -304,6 +305,32 @@ class DefaultSockJsServiceTests extends AbstractHttpRequestTests { verifyNoMoreInteractions(this.xhrSendHandler); } + @Test + void handleTransportRequestXhrSendWithSameRemoteAddress() { + String sockJsPath = sessionUrlPrefix + "xhr"; + setRequest("POST", sockJsPrefix + sockJsPath); + this.service.handleRequest(this.request, this.response, sockJsPath, this.wsHandler); + + // session created + assertThat(this.servletResponse.getStatus()).isEqualTo(200); + verify(this.xhrHandler).handleRequest(this.request, this.response, this.wsHandler, this.session); + + this.session.setRemoteAddress(new InetSocketAddress("0.0.0.0.1", 54001)); + this.servletRequest.setRemoteAddr("0.0.0.0.1"); + this.servletRequest.setRemotePort(54002); // port can vary + + resetResponse(); + reset(this.xhrSendHandler); + given(this.xhrSendHandler.checkSessionType(this.session)).willReturn(true); + + sockJsPath = sessionUrlPrefix + "xhr_send"; + setRequest("POST", sockJsPrefix + sockJsPath); + this.service.handleRequest(this.request, this.response, sockJsPath, this.wsHandler); + + assertThat(this.servletResponse.getStatus()).isEqualTo(200); + verify(this.xhrSendHandler).handleRequest(this.request, this.response, this.wsHandler, this.session); + } + @Test void handleTransportRequestWebsocket() { TransportHandlingSockJsService wsService = new TransportHandlingSockJsService(