feat(core): add WebSocket over HTTP/2

This commit is contained in:
Zakaria El Orche
2026-08-13 20:18:05 +00:00
parent 3c1eb0d0df
commit f3011ffdf6
18 changed files with 877 additions and 45 deletions
@@ -55,6 +55,7 @@ public final class Http2Preface {
setting(bytes, Http2Settings.MAX_CONCURRENT_STREAMS, Http2Limits.MAX_CONCURRENT_STREAMS);
setting(bytes, Http2Settings.INITIAL_WINDOW_SIZE, Http2Limits.INITIAL_WINDOW_SIZE_LOCAL);
setting(bytes, Http2Settings.MAX_HEADER_LIST_SIZE, Http2Limits.MAX_HEADER_LIST_SIZE);
setting(bytes, Http2Settings.ENABLE_CONNECT_PROTOCOL, 1);
frame.endFrame();
return copy(bytes);
}
@@ -11,6 +11,7 @@ public final class Http2Settings {
public static final int INITIAL_WINDOW_SIZE = 0x4;
public static final int MAX_FRAME_SIZE = 0x5;
public static final int MAX_HEADER_LIST_SIZE = 0x6;
public static final int ENABLE_CONNECT_PROTOCOL = 0x8;
public static final int DEFAULT_HEADER_TABLE_SIZE = 4_096;
public static final int DEFAULT_INITIAL_WINDOW_SIZE = 65_535;
@@ -78,7 +79,7 @@ public final class Http2Settings {
private static void validate(int id, long value) {
switch (id) {
case ENABLE_PUSH -> {
case ENABLE_PUSH, ENABLE_CONNECT_PROTOCOL -> {
if (value > 1) throw Http2Exception.PROTOCOL_ERROR;
}
case INITIAL_WINDOW_SIZE -> {
@@ -1,5 +1,6 @@
package dev.relism.flash.http2;
import dev.relism.flash.http.ContentType;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.http.HttpStatus;
import dev.relism.flash.http2.frame.Http2FrameWriter;
@@ -11,7 +12,11 @@ import dev.relism.flash.http2.stream.Http2StreamTable;
import dev.relism.flash.models.Request;
import dev.relism.flash.models.RequestHandler;
import dev.relism.flash.models.Response;
import dev.relism.flash.models.ResponseStreamOutputStream;
import dev.relism.flash.transport.ConnectionContext;
import dev.relism.flash.websocket.WebSocketHandler;
import dev.relism.flash.websocket.WebSocketLoop;
import dev.relism.flash.websocket.WebSocketSession;
import java.io.IOException;
import java.util.concurrent.RejectedExecutionException;
import lombok.extern.slf4j.Slf4j;
@@ -103,10 +108,29 @@ final class Http2StreamDispatcher implements Http2Stream.ResponseSink {
Request request = stream.assembleRequest(context.remoteAddress(), context.sslSocket());
Response pooled = stream.resetResponse();
Response response = pooled;
Object routeScratch = stream.routeScratch(context.router());
if (!Http2Authority.isServed(request.header("host"), request.sslSession())) {
response.status(HttpStatus.MISDIRECTED_REQUEST);
if (stream.websocketConnect()) response.type(ContentType.NONE).streaming(output -> {});
} else if (stream.websocketConnect()) {
WebSocketHandler handler =
context.wsRouter().route(request, stream.wsRouteScratch(context.wsRouter()));
response.type(ContentType.NONE);
if (handler == null) {
response.status(HttpStatus.NOT_FOUND).streaming(output -> {});
} else {
response.streaming(
output ->
WebSocketLoop.run(
new WebSocketSession(
request.body().stream(),
new ResponseStreamOutputStream(output),
context.configuration().getWsFrameBufferSize(),
request,
false),
handler));
}
} else {
Object routeScratch = stream.routeScratch(context.router());
RequestHandler handler = context.router().route(request, routeScratch);
if (handler == null) handler = context.router().getNotFoundHandler();
try {
@@ -151,7 +151,8 @@ public final class Http2ResponseWriter implements WriteIntent, ResponseSerialize
throws IOException {
if (streamId <= 0) throw new IllegalArgumentException("streamId must be positive");
if (maxFrameSize <= 0 || availableFlowWindow < 0) {
throw new IllegalArgumentException("frame size must be positive and flow window non-negative");
throw new IllegalArgumentException(
"frame size must be positive and flow window non-negative");
}
headerBlock.reset();
output.reset();
@@ -212,7 +213,7 @@ public final class Http2ResponseWriter implements WriteIntent, ResponseSerialize
finished = true;
endStreamInBatch = true;
}
if (hasBody && availableFlowWindow > 0) {
if (hasBody && availableFlowWindow > 0 && !pushBody) {
appendData(maxFrameSize, availableFlowWindow);
}
return dataBytesInBatch;
@@ -282,7 +283,8 @@ public final class Http2ResponseWriter implements WriteIntent, ResponseSerialize
end = unknownLength ? eof : bodyRemaining == 0;
boolean trailersFollow = end && response.hasTrailers();
if (count != 0 || !trailersFollow) {
frames.beginFrame(FrameType.DATA, end && !trailersFollow ? FrameFlags.END_STREAM : 0, streamId);
frames.beginFrame(
FrameType.DATA, end && !trailersFollow ? FrameFlags.END_STREAM : 0, streamId);
output.writeBytes(relay, 0, count);
frames.endFrame();
}
@@ -12,6 +12,7 @@ public final class PseudoHeaders {
private static final int SCHEME = 2;
private static final int PATH = 4;
private static final int AUTHORITY = 8;
private static final int PROTOCOL = 16;
private final PooledSlice name = new PooledSlice();
private final PooledSlice value = new PooledSlice();
@@ -19,6 +20,7 @@ public final class PseudoHeaders {
private final PooledSlice scheme = new PooledSlice();
private final PooledSlice path = new PooledSlice();
private final PooledSlice authority = new PooledSlice();
private final PooledSlice protocol = new PooledSlice();
private final PooledSlice host = new PooledSlice();
private int present;
@@ -28,6 +30,7 @@ public final class PseudoHeaders {
scheme.reset(null, 0, 0);
path.reset(null, 0, 0);
authority.reset(null, 0, 0);
protocol.reset(null, 0, 0);
host.reset(null, 0, 0);
boolean regularSeen = false;
@@ -51,7 +54,15 @@ public final class PseudoHeaders {
if ((present & METHOD) == 0) fail(streamId, "missing :method");
boolean connect = equals(method, "CONNECT");
if (connect) {
boolean extendedConnect = (present & PROTOCOL) != 0;
if (extendedConnect) {
if (!connect) fail(streamId, ":protocol requires CONNECT");
int required = METHOD | SCHEME | PATH | AUTHORITY | PROTOCOL;
if ((present & required) != required) {
fail(streamId, "extended CONNECT missing pseudo-header");
}
if (path.length() == 0) fail(streamId, "empty :path");
} else if (connect) {
if ((present & AUTHORITY) == 0) fail(streamId, "CONNECT requires :authority");
if ((present & (SCHEME | PATH)) != 0) fail(streamId, "CONNECT forbids :scheme and :path");
} else {
@@ -94,11 +105,16 @@ public final class PseudoHeaders {
return authority;
}
public boolean websocket() {
return protocol.array() != null && equals(protocol, "websocket");
}
private void copySlice(int bit, PooledSlice source) {
if (bit == METHOD) copy(source, method);
else if (bit == SCHEME) copy(source, scheme);
else if (bit == PATH) copy(source, path);
else copy(source, authority);
else if (bit == AUTHORITY) copy(source, authority);
else copy(source, protocol);
}
private static void copy(PooledSlice source, PooledSlice target) {
@@ -110,6 +126,7 @@ public final class PseudoHeaders {
if (equals(name, ":scheme")) return SCHEME;
if (equals(name, ":path")) return PATH;
if (equals(name, ":authority")) return AUTHORITY;
if (equals(name, ":protocol")) return PROTOCOL;
return 0;
}
@@ -17,6 +17,7 @@ import dev.relism.flash.models.RequestBody;
import dev.relism.flash.models.RequestLine;
import dev.relism.flash.models.Response;
import dev.relism.flash.routing.AbstractRouter;
import dev.relism.flash.routing.AbstractWsRouter;
import dev.relism.fpr.core.ByteView;
import java.io.IOException;
import java.net.InetSocketAddress;
@@ -60,6 +61,7 @@ public final class Http2Stream
private int emptyDataFrames;
private Http2StreamTable owner;
private Object routeScratch;
private Object wsRouteScratch;
private volatile boolean dispatched;
private volatile boolean cancelled;
private boolean headersValidated;
@@ -135,9 +137,11 @@ public final class Http2Stream
throw new Http2StreamException(
id, Http2ErrorCode.PROTOCOL_ERROR, "unsupported request method");
}
if (pseudoHeaders.websocket()) method = HttpMethod.GET;
requestLine.reset(method, path, question < 0 ? null : query, protocol, headers);
requestBody.reset(http2Body, http2Body.declaredLength(), null, 0, 0);
Request assembled = Request.forParsed(request, requestLine, requestBody, remoteAddress, sslSocket);
Request assembled =
Request.forParsed(request, requestLine, requestBody, remoteAddress, sslSocket);
assembled.setTrailers(trailers);
return assembled;
}
@@ -279,6 +283,15 @@ public final class Http2Stream
return routeScratch;
}
public Object wsRouteScratch(AbstractWsRouter router) {
if (wsRouteScratch == null) wsRouteScratch = router.newScratch();
return wsRouteScratch;
}
public boolean websocketConnect() {
return pseudoHeaders.websocket();
}
public void markDispatched() {
dispatched = true;
}
@@ -0,0 +1,35 @@
package dev.relism.flash.models;
import java.io.IOException;
import java.io.OutputStream;
/** Adapts a flow-controlled response stream to APIs that write to an {@link OutputStream}. */
public final class ResponseStreamOutputStream extends OutputStream {
private final ResponseStream stream;
private final byte[] single = new byte[1];
public ResponseStreamOutputStream(ResponseStream stream) {
this.stream = stream;
}
@Override
public void write(int value) throws IOException {
single[0] = (byte) value;
stream.write(single, 0, 1);
}
@Override
public void write(byte[] bytes, int offset, int length) throws IOException {
stream.write(bytes, offset, length);
}
@Override
public void flush() throws IOException {
stream.flush();
}
@Override
public void close() throws IOException {
stream.close();
}
}
@@ -3,39 +3,44 @@ package dev.relism.flash.websocket;
import java.io.IOException;
/**
* Drives one {@link WebSocketSession}'s read loop until the session closes, dispatching frames
* responsibility is this loop; the handshake and upgrade detection live in
* {@link WebSocketUpgrade}.
* Drives one {@link WebSocketSession}'s read loop until the session closes. The handshake and
* upgrade detection live in {@link WebSocketUpgrade}.
*/
public final class WebSocketLoop {
private WebSocketLoop() {
}
private WebSocketLoop() {}
public static void run(WebSocketSession session, WebSocketHandler handler) {
handler.onOpen(session);
WebSocketFrame frame = new WebSocketFrame();
try {
while (session.isOpen()) {
if (!session.readFrame(frame)) break;
switch (frame.opcode()) {
case WebSocketFrame.OP_TEXT, WebSocketFrame.OP_BINARY
-> handler.onMessage(session, frame);
case WebSocketFrame.OP_CLOSE
-> session.closeFromPeer(frame);
case WebSocketFrame.OP_PING
-> session.sendPong(frame);
case WebSocketFrame.OP_PONG -> { /* heartbeat ack, no-op */ }
}
}
} catch (WebSocketProtocolException e) {
try { session.close(e.closeCode()); } catch (IOException ignored) { }
handler.onError(session, e);
} catch (IOException e) {
handler.onError(session, e);
} finally {
handler.onClose(session, session.closeCode());
session.forceClose();
public static void run(WebSocketSession session, WebSocketHandler handler) {
WebSocketFrame frame = new WebSocketFrame();
try {
handler.onOpen(session);
while (session.isOpen()) {
if (!session.readFrame(frame)) break;
switch (frame.opcode()) {
case WebSocketFrame.OP_TEXT, WebSocketFrame.OP_BINARY ->
handler.onMessage(session, frame);
case WebSocketFrame.OP_CLOSE -> session.closeFromPeer(frame);
case WebSocketFrame.OP_PING -> session.sendPong(frame);
case WebSocketFrame.OP_PONG -> {
// Heartbeat acknowledgement; no action is required.
}
}
}
} catch (WebSocketProtocolException failure) {
try {
session.close(failure.closeCode());
} catch (IOException ignored) {
// The peer may already have closed the transport.
}
handler.onError(session, failure);
} catch (IOException | RuntimeException failure) {
handler.onError(session, failure);
} finally {
try {
handler.onClose(session, session.closeCode());
} finally {
session.forceClose();
}
}
}
}
@@ -0,0 +1,292 @@
package dev.relism.flash.http2;
import dev.relism.flash.bytes.ByteWriter;
import dev.relism.flash.http2.frame.FrameFlags;
import dev.relism.flash.http2.frame.FrameType;
import dev.relism.flash.http2.frame.FrameWriteBuffer;
import dev.relism.flash.http2.hpack.HpackDecoder;
import dev.relism.flash.http2.hpack.HpackEncoder;
import dev.relism.flash.websocket.WebSocketFrame;
import java.io.ByteArrayOutputStream;
import java.io.Closeable;
import java.io.EOFException;
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.Socket;
import java.nio.charset.StandardCharsets;
/** Minimal RFC 8441 peer used only by the live WebSocket-over-h2 tests. */
final class H2WebSocketTestClient implements Closeable {
private static final int WINDOW = 2 * 1024 * 1024;
private final Socket socket;
private final InputStream input;
private final OutputStream output;
private final ByteArrayOutputStream responseData = new ByteArrayOutputStream();
private int connectionWindow = 65_535;
private int streamWindow = 65_535;
private int peerMaxFrame = 16_384;
private boolean connectProtocolAdvertised;
private boolean responseEnded;
H2WebSocketTestClient(String host, int port, String path) throws Exception {
socket = new Socket(host, port);
socket.setSoTimeout(5_000);
input = socket.getInputStream();
output = socket.getOutputStream();
writePreface();
awaitSettings();
writeConnect(host + ":" + port, path);
int status = awaitStatus();
if (status != 200) throw new IOException("extended CONNECT returned " + status);
}
boolean connectProtocolAdvertised() {
return connectProtocolAdvertised;
}
void sendText(String value) throws Exception {
sendWebSocketFrame(true, WebSocketFrame.OP_TEXT, value.getBytes(StandardCharsets.UTF_8), false);
}
void sendFragmentedText(String first, String second) throws Exception {
sendWebSocketFrame(
false, WebSocketFrame.OP_TEXT, first.getBytes(StandardCharsets.UTF_8), false);
sendWebSocketFrame(
true, WebSocketFrame.OP_CONTINUATION, second.getBytes(StandardCharsets.UTF_8), false);
}
void sendBinary(byte[] value) throws Exception {
sendWebSocketFrame(true, WebSocketFrame.OP_BINARY, value, false);
}
byte[] readMessage(byte expectedOpcode) throws Exception {
responseData.reset();
while (true) {
readAndHandleFrame();
byte[] bytes = responseData.toByteArray();
if (bytes.length < 2) continue;
int opcode = bytes[0] & 0x0f;
int marker = bytes[1] & 0x7f;
int headerLength;
long payloadLength;
if (marker < 126) {
headerLength = 2;
payloadLength = marker;
} else if (marker == 126) {
if (bytes.length < 4) continue;
headerLength = 4;
payloadLength = ((bytes[2] & 0xff) << 8) | (bytes[3] & 0xff);
} else {
if (bytes.length < 10) continue;
headerLength = 10;
payloadLength = 0;
for (int i = 2; i < 10; i++) payloadLength = (payloadLength << 8) | (bytes[i] & 0xffL);
}
if (payloadLength > Integer.MAX_VALUE || bytes.length < headerLength + payloadLength) {
continue;
}
if (opcode != expectedOpcode) throw new IOException("unexpected WebSocket opcode " + opcode);
byte[] payload = new byte[(int) payloadLength];
System.arraycopy(bytes, headerLength, payload, 0, payload.length);
return payload;
}
}
void closeGracefully() throws Exception {
sendWebSocketFrame(true, WebSocketFrame.OP_CLOSE, new byte[] {3, (byte) 232}, true);
while (!responseEnded) readAndHandleFrame();
}
@Override
public void close() throws IOException {
socket.close();
}
private void writePreface() throws IOException {
output.write(Http2Preface.clientPreface());
ByteWriter bytes = new ByteWriter(64);
FrameWriteBuffer frames = new FrameWriteBuffer(bytes);
frames.beginFrame(FrameType.SETTINGS, 0, 0);
bytes.writeUInt16(Http2Settings.ENABLE_PUSH);
bytes.writeUInt32(0);
bytes.writeUInt16(Http2Settings.INITIAL_WINDOW_SIZE);
bytes.writeUInt32(WINDOW);
frames.endFrame();
frames.beginFrame(FrameType.WINDOW_UPDATE, 0, 0);
bytes.writeUInt31(WINDOW - 65_535);
frames.endFrame();
output.write(bytes.array(), 0, bytes.length());
}
private void awaitSettings() throws Exception {
while (!connectProtocolAdvertised) {
WireFrame frame = readFrame();
if (frame.type == FrameType.SETTINGS.code() && (frame.flags & FrameFlags.ACK) == 0) {
for (int offset = 0; offset < frame.payload.length; offset += 6) {
int id = ((frame.payload[offset] & 0xff) << 8) | (frame.payload[offset + 1] & 0xff);
int value = readInt(frame.payload, offset + 2);
if (id == Http2Settings.ENABLE_CONNECT_PROTOCOL && value == 1) {
connectProtocolAdvertised = true;
} else if (id == Http2Settings.INITIAL_WINDOW_SIZE) {
streamWindow = value;
} else if (id == Http2Settings.MAX_FRAME_SIZE) {
peerMaxFrame = value;
}
}
writeEmpty(FrameType.SETTINGS, FrameFlags.ACK, 0);
} else {
handle(frame);
}
}
}
private void writeConnect(String authority, String path) throws IOException {
ByteWriter bytes = new ByteWriter(256);
FrameWriteBuffer frame = new FrameWriteBuffer(bytes);
frame.beginFrame(FrameType.HEADERS, FrameFlags.END_HEADERS, 1);
HpackEncoder.writeLiteralWithNameIndex(
bytes, 2, "CONNECT".getBytes(StandardCharsets.US_ASCII), false);
HpackEncoder.writeIndexed(bytes, 6);
HpackEncoder.writeLiteralWithNameIndex(
bytes, 1, authority.getBytes(StandardCharsets.US_ASCII), false);
HpackEncoder.writeLiteralWithNameIndex(
bytes, 4, path.getBytes(StandardCharsets.US_ASCII), false);
HpackEncoder.writeLiteral(
bytes,
":protocol".getBytes(StandardCharsets.US_ASCII),
"websocket".getBytes(StandardCharsets.US_ASCII));
frame.endFrame();
output.write(bytes.array(), 0, bytes.length());
}
private int awaitStatus() throws Exception {
HpackDecoder decoder = new HpackDecoder();
while (true) {
WireFrame frame = readFrame();
if (frame.type != FrameType.HEADERS.code() || frame.streamId != 1) {
handle(frame);
continue;
}
int[] status = {0};
decoder.decode(
frame.payload,
0,
frame.payload.length,
(name, value, never) -> {
if (name.length() == 7 && name.byteAt(0) == ':') {
status[0] =
(value.byteAt(0) - '0') * 100
+ (value.byteAt(1) - '0') * 10
+ value.byteAt(2)
- '0';
}
});
return status[0];
}
}
private void sendWebSocketFrame(boolean fin, byte opcode, byte[] payload, boolean endStream)
throws Exception {
byte[] encoded = maskedFrame(fin, opcode, payload);
int offset = 0;
while (offset < encoded.length) {
while (connectionWindow <= 0 || streamWindow <= 0) readAndHandleFrame();
int count =
Math.min(
encoded.length - offset,
Math.min(peerMaxFrame, Math.min(connectionWindow, streamWindow)));
writeData(encoded, offset, count, endStream && offset + count == encoded.length);
offset += count;
connectionWindow -= count;
streamWindow -= count;
}
}
private void readAndHandleFrame() throws Exception {
handle(readFrame());
}
private void handle(WireFrame frame) throws IOException {
if (frame.type == FrameType.WINDOW_UPDATE.code()) {
int increment = readInt(frame.payload, 0) & 0x7fff_ffff;
if (frame.streamId == 0) connectionWindow += increment;
else if (frame.streamId == 1) streamWindow += increment;
} else if (frame.type == FrameType.DATA.code() && frame.streamId == 1) {
responseData.write(frame.payload);
responseEnded = (frame.flags & FrameFlags.END_STREAM) != 0;
} else if (frame.type == FrameType.SETTINGS.code() && (frame.flags & FrameFlags.ACK) == 0) {
writeEmpty(FrameType.SETTINGS, FrameFlags.ACK, 0);
} else if (frame.type == FrameType.RST_STREAM.code() && frame.streamId == 1) {
throw new IOException("WebSocket stream reset with " + readInt(frame.payload, 0));
} else if (frame.type == FrameType.GOAWAY.code()) {
throw new IOException("HTTP/2 connection closed with " + readInt(frame.payload, 4));
}
}
private void writeData(byte[] payload, int offset, int length, boolean endStream)
throws IOException {
ByteWriter bytes = new ByteWriter(length + 9);
FrameWriteBuffer frame = new FrameWriteBuffer(bytes);
frame.beginFrame(FrameType.DATA, endStream ? FrameFlags.END_STREAM : 0, 1);
bytes.writeBytes(payload, offset, length);
frame.endFrame();
output.write(bytes.array(), 0, bytes.length());
}
private void writeEmpty(FrameType type, int flags, int streamId) throws IOException {
byte[] frame = {0, 0, 0, (byte) type.code(), (byte) flags, 0, 0, 0, (byte) streamId};
output.write(frame);
}
private WireFrame readFrame() throws IOException {
byte[] header = input.readNBytes(9);
if (header.length != 9) throw new EOFException("HTTP/2 connection closed between frames");
int length = ((header[0] & 0xff) << 16) | ((header[1] & 0xff) << 8) | (header[2] & 0xff);
int streamId =
((header[5] & 0x7f) << 24)
| ((header[6] & 0xff) << 16)
| ((header[7] & 0xff) << 8)
| (header[8] & 0xff);
byte[] payload = input.readNBytes(length);
if (payload.length != length) throw new EOFException("HTTP/2 frame truncated");
return new WireFrame(header[3] & 0xff, header[4] & 0xff, streamId, payload);
}
private static byte[] maskedFrame(boolean fin, byte opcode, byte[] payload) {
int lengthBytes = payload.length <= 125 ? 0 : payload.length <= 0xffff ? 2 : 8;
byte[] frame = new byte[2 + lengthBytes + 4 + payload.length];
int position = 0;
frame[position++] = (byte) ((fin ? 0x80 : 0) | opcode);
if (lengthBytes == 0) {
frame[position++] = (byte) (0x80 | payload.length);
} else if (lengthBytes == 2) {
frame[position++] = (byte) (0x80 | 126);
frame[position++] = (byte) (payload.length >>> 8);
frame[position++] = (byte) payload.length;
} else {
frame[position++] = (byte) (0x80 | 127);
long payloadLength = payload.length;
for (int shift = 56; shift >= 0; shift -= 8) {
frame[position++] = (byte) (payloadLength >>> shift);
}
}
byte[] mask = {1, 2, 3, 4};
System.arraycopy(mask, 0, frame, position, mask.length);
position += mask.length;
for (int i = 0; i < payload.length; i++) {
frame[position + i] = (byte) (payload[i] ^ mask[i & 3]);
}
return frame;
}
private static int readInt(byte[] bytes, int offset) {
return ((bytes[offset] & 0xff) << 24)
| ((bytes[offset + 1] & 0xff) << 16)
| ((bytes[offset + 2] & 0xff) << 8)
| (bytes[offset + 3] & 0xff);
}
private record WireFrame(int type, int flags, int streamId, byte[] payload) {}
}
@@ -39,8 +39,10 @@ class Http2SettingsTest {
}
@Test
void validatesEnablePushInitialWindowAndFrameSize() {
void validatesBooleanSettingsInitialWindowAndFrameSize() {
assertCode(Http2ErrorCode.PROTOCOL_ERROR, payload(Http2Settings.ENABLE_PUSH, 2));
assertCode(
Http2ErrorCode.PROTOCOL_ERROR, payload(Http2Settings.ENABLE_CONNECT_PROTOCOL, 2));
assertCode(
Http2ErrorCode.FLOW_CONTROL_ERROR, payload(Http2Settings.INITIAL_WINDOW_SIZE, 0x8000_0000));
assertCode(Http2ErrorCode.PROTOCOL_ERROR, payload(Http2Settings.MAX_FRAME_SIZE, 16_383));
@@ -0,0 +1,75 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
import dev.relism.flash.websocket.WebSocketFrame;
import dev.relism.flash.websocket.WebSocketHandler;
import dev.relism.flash.websocket.WebSocketSession;
import java.net.ServerSocket;
import java.nio.charset.StandardCharsets;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
class WebSocketOverH2Test {
private FlashApp app;
@AfterEach
void stop() {
if (app != null) app.stop().join();
}
@Test
void opensEchoesFragmentsAndCarriesAMessageLargerThanTheFlowWindow() throws Exception {
int port = freePort();
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.wsFrameBufferSize(2 * 1024 * 1024)
.build());
app.ws(
"/chat",
new WebSocketHandler() {
@Override
public void onOpen(WebSocketSession session) {}
@Override
public void onMessage(WebSocketSession session, WebSocketFrame frame) {
try {
session.echo(frame);
} catch (Exception failure) {
throw new RuntimeException(failure);
}
}
});
app.start();
try (H2WebSocketTestClient client =
new H2WebSocketTestClient("127.0.0.1", port, "/chat")) {
assertTrue(client.connectProtocolAdvertised());
client.sendFragmentedText("hel", "lo");
assertArrayEquals(
"hello".getBytes(StandardCharsets.UTF_8),
client.readMessage(WebSocketFrame.OP_TEXT));
byte[] large = new byte[Http2Limits.INITIAL_WINDOW_SIZE_LOCAL + 128 * 1024 + 17];
for (int i = 0; i < large.length; i++) large[i] = (byte) (i * 31);
client.sendBinary(large);
assertArrayEquals(large, client.readMessage(WebSocketFrame.OP_BINARY));
client.closeGracefully();
}
}
private static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -0,0 +1,143 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
import dev.relism.flash.websocket.WebSocketFrame;
import dev.relism.flash.websocket.WebSocketHandler;
import dev.relism.flash.websocket.WebSocketSession;
import java.io.ByteArrayOutputStream;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.ServerSocket;
import java.net.Socket;
import java.nio.charset.StandardCharsets;
import java.util.Base64;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
class WebSocketParityTest {
private FlashApp app;
@AfterEach
void stop() {
if (app != null) app.stop().join();
}
@Test
void oneRouteAndHandlerEchoTheSameMessageOverHttp1AndHttp2() throws Exception {
int port = freePort();
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.ws(
"/parity",
new WebSocketHandler() {
@Override
public void onOpen(WebSocketSession session) {}
@Override
public void onMessage(WebSocketSession session, WebSocketFrame frame) {
try {
session.echo(frame);
} catch (Exception failure) {
throw new RuntimeException(failure);
}
}
});
app.start();
byte[] expected = "same-handler".getBytes(StandardCharsets.UTF_8);
byte[] overHttp1 = exchangeOverHttp1(port, expected);
byte[] overHttp2;
try (H2WebSocketTestClient client =
new H2WebSocketTestClient("127.0.0.1", port, "/parity")) {
client.sendText(new String(expected, StandardCharsets.UTF_8));
overHttp2 = client.readMessage(WebSocketFrame.OP_TEXT);
client.closeGracefully();
}
assertArrayEquals(expected, overHttp1);
assertArrayEquals(overHttp1, overHttp2);
}
private static byte[] exchangeOverHttp1(int port, byte[] payload) throws Exception {
try (Socket socket = new Socket("127.0.0.1", port)) {
socket.setSoTimeout(5_000);
InputStream input = socket.getInputStream();
OutputStream output = socket.getOutputStream();
String key =
Base64.getEncoder()
.encodeToString("flash-parity-key".getBytes(StandardCharsets.US_ASCII));
String request =
"GET /parity HTTP/1.1\r\n"
+ "Host: 127.0.0.1:"
+ port
+ "\r\nUpgrade: websocket\r\n"
+ "Connection: Upgrade\r\nSec-WebSocket-Key: "
+ key
+ "\r\nSec-WebSocket-Version: 13\r\n\r\n";
output.write(request.getBytes(StandardCharsets.US_ASCII));
output.flush();
assertTrue(readHeaders(input).startsWith("HTTP/1.1 101 Switching Protocols"));
output.write(maskedFrame(WebSocketFrame.OP_TEXT, payload));
output.flush();
byte[] echoed = readServerFrame(input, WebSocketFrame.OP_TEXT);
output.write(maskedFrame(WebSocketFrame.OP_CLOSE, new byte[] {3, (byte) 232}));
output.flush();
return echoed;
}
}
private static String readHeaders(InputStream input) throws Exception {
ByteArrayOutputStream bytes = new ByteArrayOutputStream();
int previous3 = -1;
int previous2 = -1;
int previous1 = -1;
int current;
while ((current = input.read()) >= 0) {
bytes.write(current);
if (previous3 == '\r' && previous2 == '\n' && previous1 == '\r' && current == '\n') {
break;
}
previous3 = previous2;
previous2 = previous1;
previous1 = current;
}
return bytes.toString(StandardCharsets.US_ASCII);
}
private static byte[] maskedFrame(byte opcode, byte[] payload) {
byte[] encoded = new byte[6 + payload.length];
encoded[0] = (byte) (0x80 | opcode);
encoded[1] = (byte) (0x80 | payload.length);
byte[] mask = {1, 2, 3, 4};
System.arraycopy(mask, 0, encoded, 2, mask.length);
for (int i = 0; i < payload.length; i++) {
encoded[6 + i] = (byte) (payload[i] ^ mask[i & 3]);
}
return encoded;
}
private static byte[] readServerFrame(InputStream input, byte expectedOpcode) throws Exception {
byte[] header = input.readNBytes(2);
if (header.length != 2 || (header[0] & 0x0f) != expectedOpcode) {
throw new AssertionError("unexpected WebSocket response frame");
}
int length = header[1] & 0x7f;
return input.readNBytes(length);
}
private static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -42,6 +42,28 @@ class PseudoHeaderValidationTest {
rejects(":method", "CONNECT", ":scheme", "https", ":authority", "example.com:443");
}
@Test
void validatesExtendedConnectShape() {
assertDoesNotThrow(
() ->
validate(
":method", "CONNECT",
":protocol", "websocket",
":scheme", "https",
":path", "/chat",
":authority", "example.com"));
rejects(
":method", "GET",
":protocol", "websocket",
":scheme", "https",
":path", "/chat",
":authority", "example.com");
rejects(
":method", "CONNECT",
":protocol", "websocket",
":authority", "example.com");
}
@Test
void rejectsUppercaseForbiddenAndInvalidTeFields() {
rejects(validWith("X-Test", "1"));
@@ -0,0 +1,92 @@
package dev.relism.flash.websocket;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertSame;
import static org.junit.jupiter.api.Assertions.assertThrows;
import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.InputStream;
import java.util.concurrent.atomic.AtomicInteger;
import java.util.concurrent.atomic.AtomicReference;
import org.junit.jupiter.api.Test;
class WebSocketLoopTest {
@Test
void onOpenFailureStillReportsErrorClosesAndReleasesSession() {
RuntimeException failure = new RuntimeException("open failed");
AtomicReference<Throwable> reported = new AtomicReference<>();
AtomicInteger closes = new AtomicInteger();
WebSocketSession session =
new WebSocketSession(
new ByteArrayInputStream(new byte[0]), new ByteArrayOutputStream(), 128);
WebSocketLoop.run(
session,
new WebSocketHandler() {
@Override
public void onOpen(WebSocketSession opened) {
throw failure;
}
@Override
public void onMessage(WebSocketSession opened, WebSocketFrame frame) {}
@Override
public void onError(WebSocketSession opened, Throwable error) {
reported.set(error);
}
@Override
public void onClose(WebSocketSession opened, int code) {
closes.incrementAndGet();
}
});
assertSame(failure, reported.get());
assertEquals(1, closes.get());
assertFalse(session.isOpen());
}
@Test
void onCloseFailureCannotPreventTransportRelease() {
AtomicInteger inputCloses = new AtomicInteger();
WebSocketSession session =
new WebSocketSession(
new InputStream() {
@Override
public int read() {
return -1;
}
@Override
public void close() {
inputCloses.incrementAndGet();
}
},
new ByteArrayOutputStream(),
128);
assertThrows(
RuntimeException.class,
() ->
WebSocketLoop.run(
session,
new WebSocketHandler() {
@Override
public void onOpen(WebSocketSession opened) {}
@Override
public void onMessage(WebSocketSession opened, WebSocketFrame frame) {}
@Override
public void onClose(WebSocketSession opened, int code) {
throw new RuntimeException("close failed");
}
}));
assertEquals(1, inputCloses.get());
assertFalse(session.isOpen());
}
}