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 lines = new ArrayList<>(); String line; while ((line = in.readLine()) != null && !line.isEmpty()) { lines.add(line); } return lines.isEmpty() ? "" : lines.get(0).trim(); } } }