77 lines
3.2 KiB
Java
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();
|
|
}
|
|
}
|
|
}
|