Files
spring-messaging-demo/sse-websocket/src/test/java/com/ankurm/ssews/WebSocketHandshakeOriginTest.java
T
2026-09-04 00:52:18 +05:30

77 lines
3.2 KiB
Java

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.
*
* <p>This matters because the WebSocket handshake is a plain HTTP GET, and the browser's
* same-origin policy does <em>not</em> 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<String> lines = new ArrayList<>();
String line;
while ((line = in.readLine()) != null && !line.isEmpty()) {
lines.add(line);
}
return lines.isEmpty() ? "<no response>" : lines.get(0).trim();
}
}
}