feat(core): WS client-mode masking, zero-copy header iteration, shared I/O relay buffer #1

Merged
Relism merged 1 commits from feature/core/ws-client-mode-io-reuse into master 2026-08-09 20:38:20 +00:00
6 changed files with 287 additions and 14 deletions
Showing only changes of commit a037456634 - Show all commits
@@ -34,9 +34,14 @@ import java.util.concurrent.atomic.AtomicReference;
* Pure I/O transport layer. Owns the {@link ServerSocket}, the virtual-thread * Pure I/O transport layer. Owns the {@link ServerSocket}, the virtual-thread
* executor, and the keep-alive accept loop. Routing is delegated to HTTP and WS routers. * executor, and the keep-alive accept loop. Routing is delegated to HTTP and WS routers.
* *
* <h3>Allocation model (unchanged)</h3> * <h3>Allocation model</h3>
* <ul> * <ul>
* <li>{@code LONG_BUF} (20 bytes) is the only {@link ThreadLocal} kept here.</li> * <li>{@code LONG_BUF} (20 bytes) and {@code STREAM_RELAY_BUFFER} (8 KB, for a streaming
* {@link Response} body — see {@link #writeStreamingBody}) are the only {@link ThreadLocal}s
* kept here. Both are per-connection, not per-request: one virtual thread runs a
* connection's whole keep-alive request loop (see {@link #process}), so a handler that
* streams a large response on every request allocates its relay buffer once per
* connection, not once per request.</li>
* <li>WS handshake SHA-1: {@link ThreadLocal}&lt;{@link MessageDigest}&gt; — one per * <li>WS handshake SHA-1: {@link ThreadLocal}&lt;{@link MessageDigest}&gt; — one per
* accept thread (there are now {@code ACCEPT_THREADS} of them, not one).</li> * accept thread (there are now {@code ACCEPT_THREADS} of them, not one).</li>
* </ul> * </ul>
@@ -122,6 +127,19 @@ class HttpServer implements ServerHandle {
private static final ThreadLocal<byte[]> LONG_BUF = ThreadLocal.withInitial(() -> new byte[20]); private static final ThreadLocal<byte[]> LONG_BUF = ThreadLocal.withInitial(() -> new byte[20]);
/**
* Relay buffer for copying a streaming {@link Response} body to the client — shared by
* {@link #writeStreamingBody}'s non-chunked path and {@link #writeChunked}, so both draw
* from the same reused array instead of each allocating its own {@code byte[8192]} (the
* non-chunked path previously relied on {@link InputStream#transferTo}, which allocates
* internally on every call). Sized to match the pre-existing behavior this replaces, not
* newly tuned — not exposed as a {@link FlashConfiguration} tunable since nothing here
* needed one before.
*/
private static final int STREAM_RELAY_BUFFER_SIZE = 8192;
private static final ThreadLocal<byte[]> STREAM_RELAY_BUFFER =
ThreadLocal.withInitial(() -> new byte[STREAM_RELAY_BUFFER_SIZE]);
private static final int SHA1_LEN = 20; private static final int SHA1_LEN = 20;
private static final int WS_ACCEPT_LEN = 28; private static final int WS_ACCEPT_LEN = 28;
@@ -248,7 +266,7 @@ class HttpServer implements ServerHandle {
out.flush(); out.flush();
request.drain(); request.drain();
WebSocketSession session = new WebSocketSession( WebSocketSession session = new WebSocketSession(
in, rawOut, configuration.getWsFrameBufferSize()); in, rawOut, configuration.getWsFrameBufferSize(), request, false);
runWsLoop(session, wsHandler); runWsLoop(session, wsHandler);
return; return;
} }
@@ -420,7 +438,7 @@ class HttpServer implements ServerHandle {
out.write(CRLF); out.write(CRLF);
out.write(keepAlive ? CONNECTION_KEEPALIVE : CONNECTION_CLOSE); out.write(keepAlive ? CONNECTION_KEEPALIVE : CONNECTION_CLOSE);
out.write(CRLF); out.write(CRLF);
response.getStream().transferTo(out); relay(response.getStream(), out);
} else { } else {
out.write(TRANSFER_CHUNKED); out.write(TRANSFER_CHUNKED);
out.write(keepAlive ? CONNECTION_KEEPALIVE : CONNECTION_CLOSE); out.write(keepAlive ? CONNECTION_KEEPALIVE : CONNECTION_CLOSE);
@@ -429,6 +447,18 @@ class HttpServer implements ServerHandle {
} }
} }
/**
* Copies {@code in} to {@code out} until EOF, same contract as {@link InputStream#transferTo}
* — but via {@link #STREAM_RELAY_BUFFER} instead of a fresh {@code byte[]} per call, which is
* what {@code transferTo}'s own (JDK-internal) implementation would otherwise allocate on
* every streamed response.
*/
private static void relay(InputStream in, OutputStream out) throws IOException {
byte[] buf = STREAM_RELAY_BUFFER.get();
int n;
while ((n = in.read(buf)) > 0) out.write(buf, 0, n);
}
private static void writeStatusPhrase(OutputStream out, int statusCode) throws IOException { private static void writeStatusPhrase(OutputStream out, int statusCode) throws IOException {
byte[] phrase = HttpStatus.bytesForCode(statusCode); byte[] phrase = HttpStatus.bytesForCode(statusCode);
if (phrase != null) out.write(phrase); if (phrase != null) out.write(phrase);
@@ -447,7 +477,7 @@ class HttpServer implements ServerHandle {
} }
private static void writeChunked(OutputStream out, InputStream stream) throws IOException { private static void writeChunked(OutputStream out, InputStream stream) throws IOException {
byte[] buf = new byte[8192]; byte[] buf = STREAM_RELAY_BUFFER.get();
int n; int n;
while ((n = stream.read(buf)) > 0) { while ((n = stream.read(buf)) > 0) {
writeHex(out, n); writeHex(out, n);
@@ -36,6 +36,13 @@ public class HeaderMap {
private int sectionStart; private int sectionStart;
private int sectionEnd; private int sectionEnd;
// Lazily created, then reused for the life of this HeaderMap (i.e. the connection —
// see the class javadoc) across every #forEach call and every header within a call.
// Same idiom as #view's per-call anonymous ByteView, just amortized to zero allocations
// instead of two per header: the slices are repositioned in place, not reallocated.
private Slice nameSlice;
private Slice valueSlice;
/** Resets this map to the header section {@code buffer[sectionStart, sectionEnd)}. */ /** Resets this map to the header section {@code buffer[sectionStart, sectionEnd)}. */
public void reset(byte[] buffer, int sectionStart, int sectionEnd) { public void reset(byte[] buffer, int sectionStart, int sectionEnd) {
this.buffer = buffer; this.buffer = buffer;
@@ -43,6 +50,62 @@ public class HeaderMap {
this.sectionEnd = sectionEnd; this.sectionEnd = sectionEnd;
} }
/**
* Visits every header in declaration order without allocating — no per-header {@code
* String}/{@link ByteView}/list-entry object, unlike {@link #all()}. {@code name}/{@code
* value} are the same two {@link ByteView} instances on every call, repositioned in place;
* they are valid only for the duration of that single {@link HeaderConsumer#accept} call —
* same "do not retain past the handler" rule as {@link #view}, just per-invocation instead
* of per-request. Prefer a non-capturing or field-reusing {@link HeaderConsumer} (see its
* javadoc) if the call site itself needs to stay allocation-free too.
*
* <p>Exists for callers that must handle an open-ended set of header names — e.g. a reverse
* proxy forwarding whatever the client sent — where {@link #first}/{@link #all}'s per-name
* lookup isn't usable because the set of names isn't known upfront.
*/
public void forEach(HeaderConsumer consumer) {
if (buffer == null) return;
if (nameSlice == null) {
nameSlice = new Slice();
valueSlice = new Slice();
}
int i = sectionStart;
while (i < sectionEnd) {
int lineEnd = findCR(i);
int colon = findColon(i, lineEnd);
if (colon != -1) {
int vs = skipSpaces(colon + 1, lineEnd);
nameSlice.start = i;
nameSlice.len = colon - i;
valueSlice.start = vs;
valueSlice.len = lineEnd - vs;
consumer.accept(nameSlice, valueSlice);
}
i = lineEnd + 2;
}
}
/**
* Callback for {@link #forEach}. Implement with a reusable, field-holding instance (reset
* before each {@code forEach} call) rather than a capturing lambda if the call site itself
* needs to be allocation-free too — a capturing lambda is its own per-call allocation, same
* as anywhere else on a hot path (see {@code docs/CODE-STYLE.md} in the Pathway project for
* the idiom this mirrors).
*/
@FunctionalInterface
public interface HeaderConsumer {
void accept(ByteView name, ByteView value);
}
/** Mutable zero-copy slice into {@link #buffer} — see {@link #forEach}. */
private final class Slice implements ByteView {
int start;
int len;
@Override public int length() { return len; }
@Override public byte byteAt(int i) { return buffer[start + i]; }
}
/** Returns the first value of header {@code name} (case-insensitive), or {@code null}. */ /** Returns the first value of header {@code name} (case-insensitive), or {@code null}. */
public String first(String name) { public String first(String name) {
long r = findFirst(name); long r = findFirst(name);
@@ -1,9 +1,12 @@
package dev.relism.flash.websocket; package dev.relism.flash.websocket;
import dev.relism.flash.models.Request;
import java.io.EOFException; import java.io.EOFException;
import java.io.IOException; import java.io.IOException;
import java.io.InputStream; import java.io.InputStream;
import java.io.OutputStream; import java.io.OutputStream;
import java.util.concurrent.ThreadLocalRandom;
import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicBoolean;
/** /**
@@ -39,21 +42,42 @@ public final class WebSocketSession {
private final InputStream in; private final InputStream in;
private final OutputStream out; private final OutputStream out;
private final byte[] readBuf; private final byte[] readBuf;
private final Request request;
private final boolean maskOutgoing;
private final AtomicBoolean open = new AtomicBoolean(true); private final AtomicBoolean open = new AtomicBoolean(true);
private int closeCode = 1000; private int closeCode = 1000;
private final byte[] hdrScratch = new byte[10]; /** 1 opcode byte + up to 8 extended-length bytes + up to 4 mask-key bytes (masked mode only). */
private final byte[] hdrScratch = new byte[14];
public WebSocketSession(InputStream in, OutputStream out, int bufferSize) { public WebSocketSession(InputStream in, OutputStream out, int bufferSize) {
this(in, out, bufferSize, null, false);
}
/**
* @param request the HTTP request that upgraded to this session, or {@code null} if the
* caller has no use for it (e.g. a session opened as a WS client rather
* than accepted as a WS server). Stored as-is, no copy.
* @param maskOutgoing {@code true} if this session is acting as a WS <em>client</em> — RFC 6455
* requires client-to-server frames to be masked, unlike the server-to-client
* direction {@link #writeFrame} originally only supported. See {@link
* #writeFrame} for how masking is applied without allocating.
*/
public WebSocketSession(InputStream in, OutputStream out, int bufferSize, Request request, boolean maskOutgoing) {
this.in = in; this.in = in;
this.out = out; this.out = out;
this.readBuf = new byte[bufferSize]; this.readBuf = new byte[bufferSize];
this.request = request;
this.maskOutgoing = maskOutgoing;
} }
public boolean isOpen() { return open.get(); } public boolean isOpen() { return open.get(); }
public int closeCode() { return closeCode; } public int closeCode() { return closeCode; }
/** The request that upgraded this connection, or {@code null} — see the 4-arg constructor. */
public Request request() { return request; }
// ── Public send API ──────────────────────────────────────────────────── // ── Public send API ────────────────────────────────────────────────────
public void sendText(byte[] utf8, int off, int len) throws IOException { public void sendText(byte[] utf8, int off, int len) throws IOException {
@@ -146,8 +170,9 @@ public final class WebSocketSession {
// ── Private ──────────────────────────────────────────────────────────── // ── Private ────────────────────────────────────────────────────────────
/** /**
* Encodes the WS frame header into {@link #hdrScratch} (at most 10 bytes), * Encodes the WS frame header into {@link #hdrScratch} (at most 14 bytes: 1 opcode + up to 8
* then writes header + payload in two bulk calls to the raw socket stream. * extended-length + up to 4 mask-key), then writes header + payload in two bulk calls to the
* raw socket stream.
* *
* <p>No {@code flush()} — {@code out} is the unbuffered socket {@link OutputStream} * <p>No {@code flush()} — {@code out} is the unbuffered socket {@link OutputStream}
* (see {@code HttpServer#process}). Each {@code write()} lands directly in the * (see {@code HttpServer#process}). Each {@code write()} lands directly in the
@@ -156,19 +181,28 @@ public final class WebSocketSession {
* (header then payload) will be merged into a single TCP segment by the kernel * (header then payload) will be merged into a single TCP segment by the kernel
* because they arrive faster than the ACK from the peer — exactly the coalescing * because they arrive faster than the ACK from the peer — exactly the coalescing
* we want, at zero cost. * we want, at zero cost.
*
* <p><b>{@link #maskOutgoing} (client mode):</b> RFC 6455 requires every client-to-server frame
* to be masked. The mask key is generated into {@link #hdrScratch} (no new allocation — same
* fixed field every frame reuses) and the payload is masked <em>in place</em> via {@link
* #unmaskInPlace} — XOR is its own inverse, so the exact routine {@link #readFrame} already
* uses to unmask an inbound payload masks an outbound one too, with no separate code path and
* no copy. This mutates the caller's {@code payload} array as a side effect: callers using
* masked mode must not reuse that buffer expecting it unchanged after the call.
*/ */
private void writeFrame(byte opcode, byte[] payload, int off, int len) throws IOException { private void writeFrame(byte opcode, byte[] payload, int off, int len) throws IOException {
synchronized (out) { synchronized (out) {
int hlen = 0; int hlen = 0;
hdrScratch[hlen++] = (byte) (0x80 | opcode); hdrScratch[hlen++] = (byte) (0x80 | opcode);
int maskBit = maskOutgoing ? 0x80 : 0x00;
if (len <= 125) { if (len <= 125) {
hdrScratch[hlen++] = (byte) len; hdrScratch[hlen++] = (byte) (maskBit | len);
} else if (len <= 0xFFFF) { } else if (len <= 0xFFFF) {
hdrScratch[hlen++] = 126; hdrScratch[hlen++] = (byte) (maskBit | 126);
hdrScratch[hlen++] = (byte) ((len >> 8) & 0xFF); hdrScratch[hlen++] = (byte) ((len >> 8) & 0xFF);
hdrScratch[hlen++] = (byte) (len & 0xFF); hdrScratch[hlen++] = (byte) (len & 0xFF);
} else { } else {
hdrScratch[hlen++] = 127; hdrScratch[hlen++] = (byte) (maskBit | 127);
hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0;
hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0; hdrScratch[hlen++] = 0;
hdrScratch[hlen++] = (byte) ((len >> 24) & 0xFF); hdrScratch[hlen++] = (byte) ((len >> 24) & 0xFF);
@@ -176,6 +210,17 @@ public final class WebSocketSession {
hdrScratch[hlen++] = (byte) ((len >> 8) & 0xFF); hdrScratch[hlen++] = (byte) ((len >> 8) & 0xFF);
hdrScratch[hlen++] = (byte) (len & 0xFF); hdrScratch[hlen++] = (byte) (len & 0xFF);
} }
if (maskOutgoing) {
// ThreadLocalRandom needs no seeding/allocation per call; the four mask bytes are
// carved out of one int, never boxed.
int mask = ThreadLocalRandom.current().nextInt();
byte m0 = (byte) (mask >>> 24), m1 = (byte) (mask >>> 16), m2 = (byte) (mask >>> 8), m3 = (byte) mask;
hdrScratch[hlen++] = m0;
hdrScratch[hlen++] = m1;
hdrScratch[hlen++] = m2;
hdrScratch[hlen++] = m3;
unmaskInPlace(payload, off, len, m0, m1, m2, m3);
}
out.write(hdrScratch, 0, hlen); out.write(hdrScratch, 0, hlen);
out.write(payload, off, len); out.write(payload, off, len);
// No flush — TCP_NODELAY handles delivery. See Javadoc above. // No flush — TCP_NODELAY handles delivery. See Javadoc above.
@@ -228,4 +228,37 @@ class HttpServerTest {
assertTrue(second.contains("Connection: keep-alive")); assertTrue(second.contains("Connection: keep-alive"));
} }
} }
/**
* Regression guard for the shared {@code STREAM_RELAY_BUFFER}: both the non-chunked
* ({@code /api/stream}) and chunked ({@code /api/chunked-out}) streaming paths reuse the same
* per-connection buffer now — sending one of each back to back on one connection must not
* leave either response corrupted by the other reusing the array mid-transfer.
*/
@Test
void testKeepAlive_streamedAndChunkedResponsesOnSameConnectionDontCorruptEachOther() throws Exception {
String streamReq = "GET /api/stream HTTP/1.1\r\nHost: localhost\r\n\r\n";
String chunkedReq = "GET /api/chunked-out HTTP/1.1\r\nHost: localhost\r\n\r\n";
try (Socket socket = new Socket("127.0.0.1", port);
OutputStream out = socket.getOutputStream();
InputStream in = socket.getInputStream()) {
socket.setSoTimeout(SOCKET_TIMEOUT_MS);
out.write(streamReq.getBytes(StandardCharsets.UTF_8));
out.write(chunkedReq.getBytes(StandardCharsets.UTF_8));
out.write(streamReq.getBytes(StandardCharsets.UTF_8));
out.flush();
String first = readOneResponse(in);
String second = readOneResponse(in);
String third = readOneResponse(in);
assertTrue(first.contains("Content-Length: 23"));
assertTrue(first.endsWith("streaming response body"));
assertTrue(second.contains("Transfer-Encoding: chunked"));
assertTrue(second.endsWith("streaming response body"));
assertTrue(third.contains("Content-Length: 23"));
assertTrue(third.endsWith("streaming response body"));
}
}
} }
@@ -90,4 +90,41 @@ class HeaderMapTest {
assertTrue(map.all("Host").isEmpty()); assertTrue(map.all("Host").isEmpty());
assertTrue(map.all().isEmpty()); assertTrue(map.all().isEmpty());
} }
// --- forEach ---
@Test
void forEach_visitsEveryHeaderInDeclarationOrder() {
HeaderMap map = parse("Host: localhost", "Accept: text/plain", "Cookie: a=1");
List<String> seen = new java.util.ArrayList<>();
map.forEach((name, value) -> seen.add(toStr(name) + "=" + toStr(value)));
assertEquals(List.of("Host=localhost", "Accept=text/plain", "Cookie=a=1"), seen);
}
@Test
void forEach_emptyMap_neverInvokesConsumer() {
HeaderMap map = new HeaderMap();
map.forEach((name, value) -> fail("must not be called on an empty map"));
}
@Test
void forEach_reusesTheSameTwoViewInstancesAcrossEveryHeader() {
// The zero-allocation contract: forEach must reposition two ByteViews in place, not
// allocate a fresh pair per header — same instances across all three calls here.
HeaderMap map = parse("A: 1", "B: 2", "C: 3");
List<ByteView> names = new java.util.ArrayList<>();
List<ByteView> values = new java.util.ArrayList<>();
map.forEach((name, value) -> { names.add(name); values.add(value); });
assertSame(names.get(0), names.get(1));
assertSame(names.get(1), names.get(2));
assertSame(values.get(0), values.get(1));
assertSame(values.get(1), values.get(2));
}
private static String toStr(ByteView v) {
byte[] b = new byte[v.length()];
for (int i = 0; i < b.length; i++) b[i] = v.byteAt(i);
return new String(b, StandardCharsets.UTF_8);
}
} }
@@ -1,14 +1,79 @@
package dev.relism.flash.websocket; package dev.relism.flash.websocket;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.models.HeaderMap;
import dev.relism.flash.models.Request;
import dev.relism.flash.models.RequestLine;
import dev.relism.fpr.core.ByteView;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
import java.io.ByteArrayInputStream; import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream; import java.io.ByteArrayOutputStream;
import java.nio.charset.StandardCharsets;
import static org.junit.jupiter.api.Assertions.*; import static org.junit.jupiter.api.Assertions.*;
class WebSocketSessionTest { class WebSocketSessionTest {
private static ByteView viewOf(String s) {
byte[] bytes = s.getBytes(StandardCharsets.UTF_8);
return new ByteView() {
public int length() { return bytes.length; }
public byte byteAt(int idx) { return bytes[idx]; }
};
}
@Test
void request_returnsWhatWasPassedToConstructor() {
RequestLine line = new RequestLine(HttpMethod.GET, viewOf("/chat"), null, viewOf("HTTP/1.1"), new HeaderMap());
Request req = new Request(line, new byte[0]);
WebSocketSession session = new WebSocketSession(
new ByteArrayInputStream(new byte[0]), new ByteArrayOutputStream(), 64, req, false);
assertSame(req, session.request());
}
@Test
void request_defaultsToNullOnThreeArgConstructor() {
WebSocketSession session = new WebSocketSession(new ByteArrayInputStream(new byte[0]), new ByteArrayOutputStream(), 64);
assertNull(session.request());
}
@Test
void sendText_masksWhenActingAsClient() throws Exception {
ByteArrayOutputStream out = new ByteArrayOutputStream();
WebSocketSession session = new WebSocketSession(
new ByteArrayInputStream(new byte[0]), out, 64, null, true);
byte[] payload = "hi".getBytes();
session.sendText(payload, 0, payload.length);
byte[] bytes = out.toByteArray();
assertEquals((byte) 0x81, bytes[0]); // FIN + TEXT
assertEquals((byte) (0x80 | 2), bytes[1]); // masked bit + length 2
byte m0 = bytes[2], m1 = bytes[3], m2 = bytes[4], m3 = bytes[5];
assertEquals((byte) ('h' ^ m0), bytes[6]);
assertEquals((byte) ('i' ^ m1), bytes[7]);
// The caller's buffer is mutated in place by the mask (documented, zero-copy tradeoff).
assertEquals((byte) ('h' ^ m0), payload[0]);
}
@Test
void sendText_doesNotMaskWhenActingAsServer() throws Exception {
ByteArrayOutputStream out = new ByteArrayOutputStream();
WebSocketSession session = new WebSocketSession(new ByteArrayInputStream(new byte[0]), out, 64);
byte[] payload = "hi".getBytes();
session.sendText(payload, 0, payload.length);
byte[] bytes = out.toByteArray();
assertEquals((byte) 0x81, bytes[0]);
assertEquals((byte) 2, bytes[1]); // no masked bit
assertEquals('h', bytes[2]);
assertEquals('i', bytes[3]);
}
@Test @Test
void close_setsClosedAndWritesFrame() throws Exception { void close_setsClosedAndWritesFrame() throws Exception {
ByteArrayOutputStream out = new ByteArrayOutputStream(); ByteArrayOutputStream out = new ByteArrayOutputStream();