package com.ankurm.ssews; import java.io.BufferedReader; import java.io.InputStreamReader; import java.io.OutputStream; import java.net.Socket; import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.List; import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.web.server.LocalServerPort; import static org.assertj.core.api.Assertions.assertThat; /** * What {@code registry.addEndpoint("/ws")} with no {@code setAllowedOrigins} actually permits. * *
This matters because the WebSocket handshake is a plain HTTP GET, and the browser's
* same-origin policy does not apply to it: a page on any site can open a WebSocket to
* your server, with the user's cookies attached, unless the server checks {@code Origin}. Spring
* checks it by default. A non-browser client sends no {@code Origin} at all and is allowed
* through, which is why the Java tests in this module connect without ceremony.
*/
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT,
properties = "dashboard.broadcast.enabled=false")
class WebSocketHandshakeOriginTest {
@LocalServerPort
int port;
@Test
void originIsCheckedByDefault() throws Exception {
System.out.println("=== GET /ws handshake, varying only the Origin header ===");
String noOrigin = handshake(null);
String sameOrigin = handshake("http://localhost:" + port);
String foreignOrigin = handshake("https://evil.example.com");
System.out.println("no Origin header -> " + noOrigin);
System.out.println("Origin: http://localhost -> " + sameOrigin);
System.out.println("Origin: https://evil... -> " + foreignOrigin);
assertThat(noOrigin).startsWith("HTTP/1.1 101");
assertThat(sameOrigin).startsWith("HTTP/1.1 101");
assertThat(foreignOrigin).startsWith("HTTP/1.1 403");
}
private String handshake(String origin) throws Exception {
try (Socket socket = new Socket("127.0.0.1", port)) {
socket.setSoTimeout(10_000);
StringBuilder req = new StringBuilder()
.append("GET /ws HTTP/1.1\r\n")
.append("Host: localhost:").append(port).append("\r\n")
.append("Upgrade: websocket\r\n")
.append("Connection: Upgrade\r\n")
.append("Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n")
.append("Sec-WebSocket-Version: 13\r\n");
if (origin != null) {
req.append("Origin: ").append(origin).append("\r\n");
}
req.append("\r\n");
OutputStream out = socket.getOutputStream();
out.write(req.toString().getBytes(StandardCharsets.UTF_8));
out.flush();
BufferedReader in = new BufferedReader(
new InputStreamReader(socket.getInputStream(), StandardCharsets.UTF_8));
List