Polish 'Make livereload websocket headers case insensitive'

See gh-26813

Closes gh-26813
This commit is contained in:
Phillip Webb
2021-06-15 16:48:04 -07:00
parent 8755512719
commit 5ca687c9a6
2 changed files with 113 additions and 10 deletions

View File

@@ -19,13 +19,27 @@ package org.springframework.boot.devtools.livereload;
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.InetAddress;
import java.net.InetSocketAddress;
import java.net.URI;
import java.net.UnknownHostException;
import java.time.Duration;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.Callable;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;
import java.util.function.Function;
import java.util.stream.Collectors;
import javax.websocket.ClientEndpointConfig;
import javax.websocket.ClientEndpointConfig.Configurator;
import javax.websocket.Endpoint;
import javax.websocket.HandshakeResponse;
import javax.websocket.WebSocketContainer;
import org.apache.tomcat.websocket.WsWebSocketContainer;
import org.awaitility.Awaitility;
@@ -34,13 +48,20 @@ import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.http.HttpHeaders;
import org.springframework.util.concurrent.ListenableFuture;
import org.springframework.web.client.RestTemplate;
import org.springframework.web.socket.CloseStatus;
import org.springframework.web.socket.PingMessage;
import org.springframework.web.socket.PongMessage;
import org.springframework.web.socket.TextMessage;
import org.springframework.web.socket.WebSocketExtension;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.WebSocketMessage;
import org.springframework.web.socket.WebSocketSession;
import org.springframework.web.socket.adapter.standard.StandardWebSocketHandlerAdapter;
import org.springframework.web.socket.adapter.standard.StandardWebSocketSession;
import org.springframework.web.socket.adapter.standard.WebSocketToStandardExtensionAdapter;
import org.springframework.web.socket.client.WebSocketClient;
import org.springframework.web.socket.client.standard.StandardWebSocketClient;
import org.springframework.web.socket.handler.TextWebSocketHandler;
@@ -94,7 +115,16 @@ class LiveReloadServerTests {
(msgs) -> msgs.size() == 2);
assertThat(messages.get(0)).contains("http://livereload.com/protocols/official-7");
assertThat(messages.get(1)).contains("command\":\"reload\"");
}
@Test // gh-26813
void triggerReloadWithUppercaseHeaders() throws Exception {
LiveReloadWebSocketHandler handler = connect(UppercaseWebSocketClient::new);
this.server.triggerReload();
List<String> messages = await().atMost(Duration.ofSeconds(10)).until(handler::getMessages,
(msgs) -> msgs.size() == 2);
assertThat(messages.get(0)).contains("http://livereload.com/protocols/official-7");
assertThat(messages.get(1)).contains("command\":\"reload\"");
}
@Test
@@ -126,7 +156,13 @@ class LiveReloadServerTests {
}
private LiveReloadWebSocketHandler connect() throws Exception {
WebSocketClient client = new StandardWebSocketClient(new WsWebSocketContainer());
return connect(StandardWebSocketClient::new);
}
private LiveReloadWebSocketHandler connect(Function<WebSocketContainer, WebSocketClient> clientFactory)
throws Exception {
WsWebSocketContainer webSocketContainer = new WsWebSocketContainer();
WebSocketClient client = clientFactory.apply(webSocketContainer);
LiveReloadWebSocketHandler handler = new LiveReloadWebSocketHandler();
client.doHandshake(handler, "ws://localhost:" + this.port + "/livereload");
handler.awaitHello();
@@ -246,4 +282,69 @@ class LiveReloadServerTests {
}
static class UppercaseWebSocketClient extends StandardWebSocketClient {
private final WebSocketContainer webSocketContainer;
UppercaseWebSocketClient(WebSocketContainer webSocketContainer) {
super(webSocketContainer);
this.webSocketContainer = webSocketContainer;
}
@Override
protected ListenableFuture<WebSocketSession> doHandshakeInternal(WebSocketHandler webSocketHandler,
HttpHeaders headers, URI uri, List<String> protocols, List<WebSocketExtension> extensions,
Map<String, Object> attributes) {
InetSocketAddress localAddress = new InetSocketAddress(getLocalHost(), uri.getPort());
InetSocketAddress remoteAddress = new InetSocketAddress(uri.getHost(), uri.getPort());
StandardWebSocketSession session = new StandardWebSocketSession(headers, attributes, localAddress,
remoteAddress);
ClientEndpointConfig endpointConfig = ClientEndpointConfig.Builder.create()
.configurator(new UppercaseWebSocketClientConfigurator(headers)).preferredSubprotocols(protocols)
.extensions(extensions.stream().map(WebSocketToStandardExtensionAdapter::new)
.collect(Collectors.toList()))
.build();
endpointConfig.getUserProperties().putAll(getUserProperties());
Endpoint endpoint = new StandardWebSocketHandlerAdapter(webSocketHandler, session);
Callable<WebSocketSession> connectTask = () -> {
this.webSocketContainer.connectToServer(endpoint, endpointConfig, uri);
return session;
};
return getTaskExecutor().submitListenable(connectTask);
}
private InetAddress getLocalHost() {
try {
return InetAddress.getLocalHost();
}
catch (UnknownHostException ex) {
return InetAddress.getLoopbackAddress();
}
}
}
private static class UppercaseWebSocketClientConfigurator extends Configurator {
private final HttpHeaders headers;
UppercaseWebSocketClientConfigurator(HttpHeaders headers) {
this.headers = headers;
}
@Override
public void beforeRequest(Map<String, List<String>> requestHeaders) {
Map<String, List<String>> uppercaseRequestHeaders = new LinkedHashMap<>();
requestHeaders.forEach((key, value) -> uppercaseRequestHeaders.put(key.toUpperCase(), value));
requestHeaders.clear();
requestHeaders.putAll(uppercaseRequestHeaders);
requestHeaders.putAll(this.headers);
}
@Override
public void afterResponse(HandshakeResponse response) {
}
}
}