diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/config/HandlersBeanDefinitionParser.java b/spring-websocket/src/main/java/org/springframework/web/socket/config/HandlersBeanDefinitionParser.java index 3679546987..f3db3ba3ce 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/config/HandlersBeanDefinitionParser.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/config/HandlersBeanDefinitionParser.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2020 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. @@ -32,6 +32,7 @@ import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.beans.factory.xml.BeanDefinitionParser; import org.springframework.beans.factory.xml.ParserContext; import org.springframework.lang.Nullable; +import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; import org.springframework.util.xml.DomUtils; import org.springframework.web.socket.server.support.OriginHandshakeInterceptor; @@ -85,7 +86,13 @@ class HandlersBeanDefinitionParser implements BeanDefinitionParser { ManagedList interceptors = WebSocketNamespaceUtils.parseBeanSubElements(interceptElem, context); String allowedOrigins = element.getAttribute("allowed-origins"); List origins = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOrigins, ",")); - interceptors.add(new OriginHandshakeInterceptor(origins)); + String allowedOriginPatterns = element.getAttribute("allowed-origin-patterns"); + List originPatterns = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOriginPatterns, ",")); + OriginHandshakeInterceptor interceptor = new OriginHandshakeInterceptor(origins); + if (!ObjectUtils.isEmpty(originPatterns)) { + interceptor.setAllowedOriginPatterns(originPatterns); + } + interceptors.add(interceptor); strategy = new WebSocketHandlerMappingStrategy(handler, interceptors); } 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 9e9b5b4c8b..a0828f3f37 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 @@ -60,6 +60,7 @@ import org.springframework.scheduling.concurrent.ThreadPoolTaskExecutor; import org.springframework.util.Assert; import org.springframework.util.ClassUtils; import org.springframework.util.MimeTypeUtils; +import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; import org.springframework.util.xml.DomUtils; import org.springframework.web.socket.WebSocketHandler; @@ -358,7 +359,13 @@ class MessageBrokerBeanDefinitionParser implements BeanDefinitionParser { ManagedList interceptors = WebSocketNamespaceUtils.parseBeanSubElements(interceptElem, ctx); String allowedOrigins = element.getAttribute("allowed-origins"); List origins = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOrigins, ",")); - interceptors.add(new OriginHandshakeInterceptor(origins)); + String allowedOriginPatterns = element.getAttribute("allowed-origin-patterns"); + List originPatterns = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOriginPatterns, ",")); + OriginHandshakeInterceptor interceptor = new OriginHandshakeInterceptor(origins); + if (!ObjectUtils.isEmpty(originPatterns)) { + interceptor.setAllowedOriginPatterns(originPatterns); + } + interceptors.add(interceptor); ConstructorArgumentValues cargs = new ConstructorArgumentValues(); cargs.addIndexedArgumentValue(0, subProtoHandler); cargs.addIndexedArgumentValue(1, handler); 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 c826febd23..9ec317e37f 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-2018 the original author or authors. + * Copyright 2002-2020 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. @@ -105,11 +105,18 @@ abstract class WebSocketNamespaceUtils { Element interceptElem = DomUtils.getChildElementByTagName(element, "handshake-interceptors"); ManagedList interceptors = WebSocketNamespaceUtils.parseBeanSubElements(interceptElem, context); + String allowedOrigins = element.getAttribute("allowed-origins"); List origins = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOrigins, ",")); sockJsServiceDef.getPropertyValues().add("allowedOrigins", origins); + + String allowedOriginPatterns = element.getAttribute("allowed-origin-patterns"); + List originPatterns = Arrays.asList(StringUtils.tokenizeToStringArray(allowedOriginPatterns, ",")); + sockJsServiceDef.getPropertyValues().add("allowedOriginPatterns", originPatterns); + RootBeanDefinition originHandshakeInterceptor = new RootBeanDefinition(OriginHandshakeInterceptor.class); originHandshakeInterceptor.getPropertyValues().add("allowedOrigins", origins); + originHandshakeInterceptor.getPropertyValues().add("allowedOriginPatterns", originPatterns); interceptors.add(originHandshakeInterceptor); sockJsServiceDef.getPropertyValues().add("handshakeInterceptors", interceptors); diff --git a/spring-websocket/src/main/resources/org/springframework/web/socket/config/spring-websocket.xsd b/spring-websocket/src/main/resources/org/springframework/web/socket/config/spring-websocket.xsd index 923a3b9f45..1870ad7ee5 100644 --- a/spring-websocket/src/main/resources/org/springframework/web/socket/config/spring-websocket.xsd +++ b/spring-websocket/src/main/resources/org/springframework/web/socket/config/spring-websocket.xsd @@ -541,6 +541,16 @@ ]]> + + + + + @@ -739,6 +749,16 @@ ]]> + + + + + diff --git a/spring-websocket/src/test/java/org/springframework/web/socket/config/HandlersBeanDefinitionParserTests.java b/spring-websocket/src/test/java/org/springframework/web/socket/config/HandlersBeanDefinitionParserTests.java index 27024631ad..6bf407c83a 100644 --- a/spring-websocket/src/test/java/org/springframework/web/socket/config/HandlersBeanDefinitionParserTests.java +++ b/spring-websocket/src/test/java/org/springframework/web/socket/config/HandlersBeanDefinitionParserTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2019 the original author or authors. + * Copyright 2002-2020 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. @@ -226,8 +226,8 @@ public class HandlersBeanDefinitionParserTests { List interceptors = transportService.getHandshakeInterceptors(); assertThat(interceptors).extracting("class").containsExactly(OriginHandshakeInterceptor.class); assertThat(transportService.shouldSuppressCors()).isTrue(); - assertThat(transportService.getAllowedOrigins().contains("https://mydomain1.example")).isTrue(); - assertThat(transportService.getAllowedOrigins().contains("https://mydomain2.example")).isTrue(); + assertThat(transportService.getAllowedOrigins()).containsExactly("https://mydomain1.example", "https://mydomain2.example"); + assertThat(transportService.getAllowedOriginPatterns()).containsExactly("https://*.mydomain.example"); } 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 af5df574ae..fda305a34f 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 @@ -181,8 +181,8 @@ public class MessageBrokerBeanDefinitionParserTests { interceptors = defaultSockJsService.getHandshakeInterceptors(); assertThat(interceptors).extracting("class").containsExactly(FooTestInterceptor.class, BarTestInterceptor.class, OriginHandshakeInterceptor.class); - assertThat(defaultSockJsService.getAllowedOrigins().contains("https://mydomain3.com")).isTrue(); - assertThat(defaultSockJsService.getAllowedOrigins().contains("https://mydomain4.com")).isTrue(); + assertThat(defaultSockJsService.getAllowedOrigins()).containsExactly("https://mydomain3.com", "https://mydomain4.com"); + assertThat(defaultSockJsService.getAllowedOriginPatterns()).containsExactly("https://*.mydomain.com"); SimpUserRegistry userRegistry = this.appContext.getBean(SimpUserRegistry.class); assertThat(userRegistry).isNotNull(); diff --git a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-simple.xml b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-simple.xml index 66301e6be4..4f640331f4 100644 --- a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-simple.xml +++ b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-broker-simple.xml @@ -17,7 +17,9 @@ - + @@ -25,7 +27,9 @@ - + diff --git a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-handlers-sockjs-attributes.xml b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-handlers-sockjs-attributes.xml index 308de21259..fd93cd9005 100644 --- a/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-handlers-sockjs-attributes.xml +++ b/spring-websocket/src/test/resources/org/springframework/web/socket/config/websocket-config-handlers-sockjs-attributes.xml @@ -5,7 +5,7 @@ http://www.springframework.org/schema/beans https://www.springframework.org/schema/beans/spring-beans.xsd http://www.springframework.org/schema/websocket https://www.springframework.org/schema/websocket/spring-websocket.xsd"> - +