feat(core): add HTTP/2 cleartext proxy support

This commit is contained in:
Zakaria El Orche
2026-08-13 20:00:59 +00:00
parent 5755ef77fe
commit 3c1eb0d0df
27 changed files with 1593 additions and 109 deletions
@@ -88,13 +88,15 @@ public class FlashConfiguration {
*/
@Builder.Default int shutdownDrainTimeoutMs = 15_000;
/**
* Whether this server negotiates HTTP/2. When enabled, plaintext listeners recognize h2c prior
* knowledge and TLS listeners advertise {@code h2} followed by HTTP/1.1 through ALPN. Disabled by
* default until the HTTP/2 request/response path is complete.
*/
/** Whether TLS listeners advertise HTTP/2 through ALPN. */
@Builder.Default boolean http2Enabled = false;
/**
* Whether plaintext listeners accept the HTTP/2 prior-knowledge preface. This is independent
* from TLS HTTP/2 and deliberately disabled by default.
*/
@Builder.Default boolean http2CleartextEnabled = false;
/**
* Whether runtime-generated HTTP/2 header values use HPACK Huffman coding. Constants are always
* compressed once at startup; leaving this disabled avoids a per-byte encode pass on responses.
@@ -0,0 +1,87 @@
package dev.relism.flash.http;
import dev.relism.flash.models.HeaderView;
import dev.relism.fpr.core.ByteView;
/** Shared proxy policy for fields that must not cross an HTTP connection boundary. */
public final class HopByHopHeaders {
public enum Protocol {
HTTP_1_1,
HTTP_2
}
private HopByHopHeaders() {}
/** Returns whether a field may be copied to a new downstream connection. */
public static boolean shouldForward(
HeaderView source,
ByteView name,
ByteView value,
Protocol sourceProtocol,
Protocol targetProtocol) {
if (name.length() == 0 || name.byteAt(0) == ':') return false;
if (is(name, "connection")
|| is(name, "keep-alive")
|| is(name, "proxy-connection")
|| is(name, "proxy-authenticate")
|| is(name, "proxy-authorization")
|| is(name, "trailer")
|| is(name, "transfer-encoding")
|| is(name, "upgrade")) {
return false;
}
if (isConnectionListed(source, name)) return false;
if (is(name, "te")) {
return targetProtocol == Protocol.HTTP_2 && isTrimmed(value, "trailers");
}
return true;
}
private static boolean isConnectionListed(HeaderView source, ByteView fieldName) {
for (String value : source.all("connection")) {
int start = 0;
while (start < value.length()) {
int comma = value.indexOf(',', start);
int end = comma < 0 ? value.length() : comma;
while (start < end && isWhitespace(value.charAt(start))) start++;
while (end > start && isWhitespace(value.charAt(end - 1))) end--;
if (equalsAsciiIgnoreCase(fieldName, value, start, end)) return true;
start = comma < 0 ? value.length() : comma + 1;
}
}
return false;
}
private static boolean is(ByteView bytes, String expected) {
return equalsAsciiIgnoreCase(bytes, expected, 0, expected.length());
}
private static boolean isTrimmed(ByteView bytes, String expected) {
int start = 0;
int end = bytes.length();
while (start < end && isWhitespace((char) bytes.byteAt(start))) start++;
while (end > start && isWhitespace((char) bytes.byteAt(end - 1))) end--;
if (end - start != expected.length()) return false;
for (int i = 0; i < expected.length(); i++) {
if (lower(bytes.byteAt(start + i) & 0xff) != lower(expected.charAt(i))) return false;
}
return true;
}
private static boolean equalsAsciiIgnoreCase(
ByteView bytes, String expected, int expectedStart, int expectedEnd) {
if (bytes.length() != expectedEnd - expectedStart) return false;
for (int i = 0; i < bytes.length(); i++) {
if (lower(bytes.byteAt(i) & 0xff) != lower(expected.charAt(expectedStart + i))) return false;
}
return true;
}
private static int lower(int value) {
return value >= 'A' && value <= 'Z' ? value + ('a' - 'A') : value;
}
private static boolean isWhitespace(char value) {
return value == ' ' || value == '\t';
}
}
@@ -0,0 +1,92 @@
package dev.relism.flash.http.proxy;
import dev.relism.flash.http.HopByHopHeaders;
import dev.relism.flash.http.HopByHopHeaders.Protocol;
import dev.relism.flash.http2.client.Http2Client;
import dev.relism.flash.http2.client.Http2ClientResponse;
import dev.relism.flash.models.HeaderView;
import dev.relism.flash.models.Request;
import dev.relism.flash.models.Response;
import dev.relism.flash.models.SimpleHandler;
import dev.relism.fpr.core.ByteView;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.util.Objects;
/** Protocol-neutral reverse-proxy adapter backed by Flash's HTTP/2 upstream client. */
public final class HttpProxy {
private HttpProxy() {}
/** Creates a handler that preserves the incoming path, query, fields, body and trailers. */
public static SimpleHandler.FunctionalHandler toHttp2(URI upstreamOrigin, Http2Client client) {
Objects.requireNonNull(upstreamOrigin, "upstreamOrigin");
Objects.requireNonNull(client, "client");
return (request, response) -> relay(upstreamOrigin, client, request, response);
}
private static Response relay(
URI upstreamOrigin, Http2Client client, Request request, Response response) throws Exception {
byte[] body = request.body().bytes();
Protocol downstream =
request.getRequestLine().getProtocol() == null ? Protocol.HTTP_2 : Protocol.HTTP_1_1;
URI target = upstreamOrigin.resolve(rawTarget(request));
Http2ClientResponse upstream =
client.exchange(
target,
request.method(),
request.getRequestLine().getHeaders(),
body,
request.trailers());
response.status(upstream.statusCode()).body(upstream.body());
copyHeaders(upstream.headers(), Protocol.HTTP_2, downstream, response, false);
copyHeaders(upstream.trailers(), Protocol.HTTP_2, downstream, response, true);
return response;
}
private static String rawTarget(Request request) {
String path = request.path();
ByteView query = request.getRequestLine().getQuery();
if (query == null || query.length() == 0) return path;
byte[] bytes = new byte[query.length()];
for (int i = 0; i < bytes.length; i++) bytes[i] = query.byteAt(i);
return path + "?" + new String(bytes, StandardCharsets.US_ASCII);
}
private static void copyHeaders(
HeaderView source,
Protocol sourceProtocol,
Protocol targetProtocol,
Response response,
boolean trailers) {
source.forEach(
(name, value) -> {
if (!HopByHopHeaders.shouldForward(
source, name, value, sourceProtocol, targetProtocol)) return;
if (!trailers && (equalsAscii(name, "content-length") || equalsAscii(name, "content-type"))) {
if (equalsAscii(name, "content-type")) response.type(string(value));
return;
}
if (trailers) response.trailer(string(name), string(value));
else response.header(string(name), string(value));
});
}
private static String string(ByteView value) {
byte[] bytes = new byte[value.length()];
for (int i = 0; i < bytes.length; i++) bytes[i] = value.byteAt(i);
return new String(bytes, StandardCharsets.UTF_8);
}
private static boolean equalsAscii(ByteView bytes, String value) {
if (bytes.length() != value.length()) return false;
for (int i = 0; i < bytes.length(); i++) {
int left = bytes.byteAt(i) & 0xff;
int right = value.charAt(i);
if (left >= 'A' && left <= 'Z') left += 'a' - 'A';
if (right >= 'A' && right <= 'Z') right += 'a' - 'A';
if (left != right) return false;
}
return true;
}
}
@@ -0,0 +1,54 @@
package dev.relism.flash.http2;
import java.security.cert.Certificate;
import java.security.cert.CertificateParsingException;
import java.security.cert.X509Certificate;
import java.util.Collection;
import java.util.List;
import java.util.Locale;
import javax.net.ssl.SSLSession;
/** Validates a coalesced request authority against the certificate selected for its connection. */
final class Http2Authority {
private Http2Authority() {}
static boolean isServed(String authority, SSLSession session) {
if (session == null || authority == null) return true;
String host = host(authority);
try {
Certificate[] certificates = session.getLocalCertificates();
if (certificates == null || certificates.length == 0
|| !(certificates[0] instanceof X509Certificate certificate)) {
return true;
}
Collection<List<?>> names = certificate.getSubjectAlternativeNames();
if (names == null) return true;
for (List<?> name : names) {
int type = (Integer) name.get(0);
if ((type == 2 || type == 7) && matches(host, name.get(1).toString())) return true;
}
return false;
} catch (CertificateParsingException failure) {
return true;
}
}
static boolean matches(String authority, String certificateName) {
String host = host(authority).toLowerCase(Locale.ROOT);
String name = certificateName.toLowerCase(Locale.ROOT);
if (!name.startsWith("*.")) return host.equals(name);
String suffix = name.substring(1);
if (!host.endsWith(suffix)) return false;
int prefixLength = host.length() - suffix.length();
return prefixLength > 0 && host.indexOf('.') == prefixLength;
}
private static String host(String authority) {
if (authority.startsWith("[")) {
int closing = authority.indexOf(']');
return closing < 0 ? authority : authority.substring(1, closing);
}
int colon = authority.lastIndexOf(':');
return colon > 0 && authority.indexOf(':') == colon ? authority.substring(0, colon) : authority;
}
}
@@ -7,7 +7,7 @@ import dev.relism.flash.http2.frame.FrameWriteBuffer;
import java.nio.charset.StandardCharsets;
/** Byte-exact client preface and immutable server startup frames, compiled once at class load. */
final class Http2Preface {
public final class Http2Preface {
private static final byte[] CLIENT_PREFACE =
"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n".getBytes(StandardCharsets.US_ASCII);
private static final byte[] SERVER_SETTINGS = buildServerSettings();
@@ -16,6 +16,11 @@ final class Http2Preface {
private Http2Preface() {}
/** Immutable client connection preface bytes. Callers must not modify the returned array. */
public static byte[] clientPreface() {
return CLIENT_PREFACE;
}
static int clientPrefaceLength() {
return CLIENT_PREFACE.length;
}
@@ -1,6 +1,7 @@
package dev.relism.flash.http2;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.http.HttpStatus;
import dev.relism.flash.http2.frame.Http2FrameWriter;
import dev.relism.flash.http2.message.Http2ResponseWriter;
import dev.relism.flash.http2.stream.Http2FlowController;
@@ -103,17 +104,21 @@ final class Http2StreamDispatcher implements Http2Stream.ResponseSink {
Response pooled = stream.resetResponse();
Response response = pooled;
Object routeScratch = stream.routeScratch(context.router());
RequestHandler handler = context.router().route(request, routeScratch);
if (handler == null) handler = context.router().getNotFoundHandler();
try {
Object result = handler.handle(request, response);
if (result instanceof Response returned) response = returned;
else if (result != null) response.setBody(result);
} catch (Exception handlerFailure) {
Object result =
context.router().getExceptionHandler().handle(handlerFailure, request, response);
if (result instanceof Response returned) response = returned;
else if (result != null) response.setBody(result);
if (!Http2Authority.isServed(request.header("host"), request.sslSession())) {
response.status(HttpStatus.MISDIRECTED_REQUEST);
} else {
RequestHandler handler = context.router().route(request, routeScratch);
if (handler == null) handler = context.router().getNotFoundHandler();
try {
Object result = handler.handle(request, response);
if (result instanceof Response returned) response = returned;
else if (result != null) response.setBody(result);
} catch (Exception handlerFailure) {
Object result =
context.router().getExceptionHandler().handle(handlerFailure, request, response);
if (result instanceof Response returned) response = returned;
else if (result != null) response.setBody(result);
}
}
boolean pushStreaming = response.isPushStreaming();
@@ -0,0 +1,621 @@
package dev.relism.flash.http2.client;
import dev.relism.flash.bytes.ByteWriter;
import dev.relism.flash.bytes.Pairs;
import dev.relism.flash.http.HopByHopHeaders;
import dev.relism.flash.http.HopByHopHeaders.Protocol;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.http2.Http2Exception;
import dev.relism.flash.http2.Http2Limits;
import dev.relism.flash.http2.Http2Preface;
import dev.relism.flash.http2.Http2Settings;
import dev.relism.flash.http2.frame.FrameFlags;
import dev.relism.flash.http2.frame.FrameHeader;
import dev.relism.flash.http2.frame.FrameType;
import dev.relism.flash.http2.frame.FrameWriteBuffer;
import dev.relism.flash.http2.frame.Http2FrameReader;
import dev.relism.flash.http2.frame.Http2FrameWriter;
import dev.relism.flash.http2.frame.Padding;
import dev.relism.flash.http2.frame.WriteIntent;
import dev.relism.flash.http2.hpack.ContinuationAssembler;
import dev.relism.flash.http2.hpack.HpackDecoder;
import dev.relism.flash.http2.hpack.HpackEncoder;
import dev.relism.flash.models.EmptyHeaderView;
import dev.relism.flash.models.HeaderView;
import dev.relism.flash.models.MutableHeaderMap;
import dev.relism.flash.transport.BufferedByteSource;
import dev.relism.fpr.core.ByteView;
import java.io.ByteArrayOutputStream;
import java.io.Closeable;
import java.io.IOException;
import java.io.OutputStream;
import java.net.InetSocketAddress;
import java.net.Socket;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
import javax.net.ssl.SSLContext;
import javax.net.ssl.SSLParameters;
import javax.net.ssl.SSLSocket;
/**
* Small pooled HTTP/2 client for Flash proxy handlers. It intentionally exposes synchronous
* request/response exchange rather than trying to be a general-purpose client API.
*/
public final class Http2Client implements Closeable {
private static final int CONNECT_TIMEOUT_MS = 10_000;
private static final int MAX_RESPONSE_BODY_SIZE = Http2Limits.MAX_REQUEST_BODY_SIZE;
private final ConcurrentHashMap<Origin, Connection> connections = new ConcurrentHashMap<>();
private final SSLContext sslContext;
public Http2Client() {
this(null);
}
public Http2Client(SSLContext sslContext) {
this.sslContext = sslContext;
}
public Http2ClientResponse get(URI uri) throws IOException {
return exchange(
uri,
HttpMethod.GET,
EmptyHeaderView.INSTANCE,
new byte[0],
EmptyHeaderView.INSTANCE);
}
public Http2ClientResponse exchange(
URI uri, HttpMethod method, HeaderView headers, byte[] body, HeaderView trailers)
throws IOException {
Objects.requireNonNull(uri, "uri");
Objects.requireNonNull(method, "method");
Objects.requireNonNull(headers, "headers");
Objects.requireNonNull(body, "body");
Objects.requireNonNull(trailers, "trailers");
Origin origin = Origin.from(uri);
Connection connection;
try {
connection = connections.computeIfAbsent(origin, this::openUnchecked);
} catch (OpenFailure failure) {
throw failure.io;
}
try {
return connection.exchange(uri, method, headers, body, trailers);
} catch (IOException | RuntimeException failure) {
connections.remove(origin, connection);
connection.close();
throw failure;
}
}
@Override
public void close() {
for (Connection connection : connections.values()) connection.close();
connections.clear();
}
/** Number of currently pooled origin connections. */
public int pooledConnectionCount() {
return connections.size();
}
private Connection openUnchecked(Origin origin) {
try {
return new Connection(origin, sslContext);
} catch (IOException failure) {
throw new OpenFailure(failure);
}
}
private static final class Connection implements Closeable {
private final Socket socket;
private final OutputStream output;
private final Http2FrameReader reader;
private final Http2FrameWriter writer;
private final Http2Settings peerSettings = new Http2Settings();
private final HpackDecoder decoder = new HpackDecoder();
private final ContinuationAssembler headers = new ContinuationAssembler();
private final ByteWriter outgoing = new ByteWriter(16 * 1024);
private final FrameWriteBuffer frames = new FrameWriteBuffer(outgoing);
private final BufferIntent intent = new BufferIntent();
private int nextStreamId = 1;
private int connectionSendWindow = Http2Settings.DEFAULT_INITIAL_WINDOW_SIZE;
private int streamSendWindow;
private boolean headerEndStream;
private boolean closed;
Connection(Origin origin, SSLContext sslContext) throws IOException {
socket = connect(origin, sslContext);
output = socket.getOutputStream();
reader =
new Http2FrameReader(new BufferedByteSource(socket.getInputStream(), socket));
writer = new Http2FrameWriter(output::write);
writePreface();
awaitServerSettings();
}
synchronized Http2ClientResponse exchange(
URI uri, HttpMethod method, HeaderView requestHeaders, byte[] body, HeaderView trailers)
throws IOException {
if (closed) throw new IOException("HTTP/2 connection is closed");
if (nextStreamId <= 0) throw new IOException("HTTP/2 stream id space exhausted");
int streamId = nextStreamId;
nextStreamId += 2;
streamSendWindow = peerSettings.initialWindowSize();
Exchange exchange = new Exchange(streamId);
writeRequestHeaders(uri, method, requestHeaders, body.length == 0 && trailers.count() == 0,
streamId);
if (body.length != 0) writeRequestBody(exchange, body, trailers.count() == 0);
if (trailers.count() != 0) writeRequestTrailers(trailers, streamId);
while (!exchange.complete) readFrame(exchange);
return exchange.response();
}
private void writePreface() throws IOException {
output.write(Http2Preface.clientPreface());
outgoing.reset();
frames.beginFrame(FrameType.SETTINGS, 0, 0);
outgoing.writeUInt16(Http2Settings.ENABLE_PUSH);
outgoing.writeUInt32(0);
outgoing.writeUInt16(Http2Settings.INITIAL_WINDOW_SIZE);
outgoing.writeUInt32(Http2Limits.INITIAL_WINDOW_SIZE_LOCAL);
frames.endFrame();
frames.beginFrame(FrameType.WINDOW_UPDATE, 0, 0);
outgoing.writeUInt31(
Http2Limits.CONNECTION_WINDOW_SIZE_LOCAL - Http2Settings.DEFAULT_INITIAL_WINDOW_SIZE);
frames.endFrame();
writeOutgoing();
}
private void awaitServerSettings() throws IOException {
boolean received = false;
while (!received) {
FrameHeader frame = reader.readFrame();
if (frame == null) throw new IOException("server closed before SETTINGS");
try {
if (frame.type() == FrameType.SETTINGS && !FrameFlags.isAck(frame.flags())) {
applySettings(frame);
sendEmpty(FrameType.SETTINGS, FrameFlags.ACK, 0);
received = true;
} else if (frame.type() == FrameType.WINDOW_UPDATE) {
applyWindowUpdate(frame, 0);
} else if (frame.type() == FrameType.GOAWAY) {
throw new IOException("server sent GOAWAY during HTTP/2 setup");
}
} finally {
reader.consumeFrame();
}
}
}
private void writeRequestHeaders(
URI uri, HttpMethod method, HeaderView source, boolean endStream, int streamId)
throws IOException {
outgoing.reset();
frames.beginFrame(
FrameType.HEADERS,
FrameFlags.END_HEADERS | (endStream ? FrameFlags.END_STREAM : 0),
streamId);
writeMethod(method);
HpackEncoder.writeIndexed(outgoing, "https".equalsIgnoreCase(uri.getScheme()) ? 7 : 6);
writeAuthority(uri);
writePath(uri);
source.forEach(
(name, value) -> {
if (HopByHopHeaders.shouldForward(
source, name, value, Protocol.HTTP_1_1, Protocol.HTTP_2)
&& !equalsAscii(name, "host")) {
HpackEncoder.writeLiteral(outgoing, name, value);
}
});
frames.endFrame();
writeOutgoing();
}
private void writeRequestBody(Exchange exchange, byte[] body, boolean endStream)
throws IOException {
int offset = 0;
while (offset < body.length) {
while (connectionSendWindow <= 0 || streamSendWindow <= 0) readFrame(exchange);
int count =
Math.min(
body.length - offset,
Math.min(
peerSettings.maxFrameSize(),
Math.min(connectionSendWindow, streamSendWindow)));
outgoing.reset();
frames.beginFrame(
FrameType.DATA,
endStream && offset + count == body.length ? FrameFlags.END_STREAM : 0,
exchange.streamId);
outgoing.writeBytes(body, offset, count);
frames.endFrame();
writeOutgoing();
connectionSendWindow -= count;
streamSendWindow -= count;
offset += count;
}
}
private void writeRequestTrailers(HeaderView trailers, int streamId) throws IOException {
outgoing.reset();
frames.beginFrame(
FrameType.HEADERS, FrameFlags.END_HEADERS | FrameFlags.END_STREAM, streamId);
trailers.forEach(
(name, value) -> {
if (HopByHopHeaders.shouldForward(
trailers, name, value, Protocol.HTTP_1_1, Protocol.HTTP_2)) {
HpackEncoder.writeLiteral(outgoing, name, value);
}
});
frames.endFrame();
writeOutgoing();
}
private void readFrame(Exchange exchange) throws IOException {
FrameHeader frame = reader.readFrame();
if (frame == null) throw new IOException("server closed an active HTTP/2 exchange");
try {
FrameType type = frame.type();
if (type == null) return;
switch (type) {
case SETTINGS -> {
if (!FrameFlags.isAck(frame.flags())) {
applySettings(frame);
sendEmpty(FrameType.SETTINGS, FrameFlags.ACK, 0);
}
}
case WINDOW_UPDATE -> applyWindowUpdate(frame, exchange.streamId);
case PING -> {
if (!FrameFlags.isAck(frame.flags())) sendPingAck(frame);
}
case HEADERS, CONTINUATION -> receiveHeaders(frame, exchange);
case DATA -> receiveData(frame, exchange);
case RST_STREAM -> receiveReset(frame, exchange);
case GOAWAY -> throw receiveGoAway(frame);
case PUSH_PROMISE -> throw new IOException("server sent PUSH_PROMISE after ENABLE_PUSH=0");
default -> {
// PRIORITY and unknown extension semantics do not affect this single exchange.
}
}
} finally {
reader.consumeFrame();
}
}
private void receiveHeaders(FrameHeader frame, Exchange exchange) throws IOException {
if (frame.streamId() != exchange.streamId) {
throw new IOException("unexpected response stream " + frame.streamId());
}
if (frame.type() == FrameType.HEADERS) {
if (headers.isActive()) throw new IOException("interleaved response header block");
headerEndStream = FrameFlags.isEndStream(frame.flags());
long unpadded =
Padding.unpad(
frame.buffer(),
frame.payloadOffset(),
frame.length(),
FrameFlags.isPadded(frame.flags()));
int offset = Pairs.hi(unpadded);
int length = Pairs.lo(unpadded);
if (FrameFlags.hasPriority(frame.flags())) {
if (length < 5) throw new IOException("truncated response priority fields");
offset += 5;
length -= 5;
}
headers.begin(
frame.streamId(),
frame.buffer(),
offset,
length,
FrameFlags.isEndHeaders(frame.flags()));
} else {
headers.continuation(
frame.streamId(),
frame.buffer(),
frame.payloadOffset(),
frame.length(),
FrameFlags.isEndHeaders(frame.flags()));
}
if (!headers.isComplete()) return;
boolean trailers = exchange.statusCode != 0;
ResponseHeaderSink sink = new ResponseHeaderSink(exchange, trailers);
decoder.decode(headers.buffer(), 0, headers.length(), sink);
headers.reset();
sink.validate();
if (!trailers && exchange.statusCode >= 100 && exchange.statusCode < 200) {
if (headerEndStream) throw new IOException("informational response ended the stream");
exchange.statusCode = 0;
exchange.headers.reset();
return;
}
if (trailers && !headerEndStream) {
throw new IOException("response trailers did not end the stream");
}
if (headerEndStream) exchange.complete = true;
}
private void receiveData(FrameHeader frame, Exchange exchange) throws IOException {
if (frame.streamId() != exchange.streamId || exchange.statusCode == 0) {
throw new IOException("DATA received before response headers");
}
long unpadded =
Padding.unpad(
frame.buffer(),
frame.payloadOffset(),
frame.length(),
FrameFlags.isPadded(frame.flags()));
int dataOffset = Pairs.hi(unpadded);
int dataLength = Pairs.lo(unpadded);
if (exchange.body.size() > MAX_RESPONSE_BODY_SIZE - dataLength) {
throw new IOException("proxied HTTP/2 response body exceeds limit");
}
exchange.body.write(frame.buffer(), dataOffset, dataLength);
if (frame.length() != 0) {
sendWindowUpdate(0, frame.length());
sendWindowUpdate(exchange.streamId, frame.length());
}
if (FrameFlags.isEndStream(frame.flags())) exchange.complete = true;
}
private void receiveReset(FrameHeader frame, Exchange exchange) throws IOException {
if (frame.streamId() != exchange.streamId || frame.length() != 4) return;
int code = readInt(frame.buffer(), frame.payloadOffset());
throw new IOException("upstream reset HTTP/2 stream with error " + code);
}
private IOException receiveGoAway(FrameHeader frame) {
closed = true;
int code = frame.length() >= 8 ? readInt(frame.buffer(), frame.payloadOffset() + 4) : -1;
return new IOException("upstream sent GOAWAY with error " + code);
}
private void applySettings(FrameHeader frame) {
int oldWindow = peerSettings.initialWindowSize();
peerSettings.apply(frame.buffer(), frame.payloadOffset(), frame.length(), delta -> {});
streamSendWindow += peerSettings.initialWindowSize() - oldWindow;
}
private void applyWindowUpdate(FrameHeader frame, int activeStreamId) throws IOException {
if (frame.length() != 4) throw new IOException("invalid WINDOW_UPDATE length");
int increment = readInt(frame.buffer(), frame.payloadOffset()) & 0x7fff_ffff;
if (increment == 0) throw new IOException("zero WINDOW_UPDATE increment");
if (frame.streamId() == 0) connectionSendWindow = addWindow(connectionSendWindow, increment);
else if (frame.streamId() == activeStreamId) streamSendWindow = addWindow(streamSendWindow, increment);
}
private void sendPingAck(FrameHeader frame) throws IOException {
outgoing.reset();
frames.beginFrame(FrameType.PING, FrameFlags.ACK, 0);
outgoing.writeBytes(frame.buffer(), frame.payloadOffset(), frame.length());
frames.endFrame();
writeOutgoing();
}
private void sendWindowUpdate(int streamId, int increment) throws IOException {
outgoing.reset();
frames.beginFrame(FrameType.WINDOW_UPDATE, 0, streamId);
outgoing.writeUInt31(increment);
frames.endFrame();
writeOutgoing();
}
private void sendEmpty(FrameType type, int flags, int streamId) throws IOException {
outgoing.reset();
frames.beginFrame(type, flags, streamId);
frames.endFrame();
writeOutgoing();
}
private void writeOutgoing() throws IOException {
intent.reset(outgoing.array(), outgoing.length());
writer.write(intent);
}
private void writeMethod(HttpMethod method) {
if (method == HttpMethod.GET) HpackEncoder.writeIndexed(outgoing, 2);
else if (method == HttpMethod.POST) HpackEncoder.writeIndexed(outgoing, 3);
else {
byte[] value = method.name().getBytes(StandardCharsets.US_ASCII);
HpackEncoder.writeLiteralWithNameIndex(outgoing, 2, value, false);
}
}
private void writeAuthority(URI uri) {
String authority = uri.getRawAuthority();
if (authority == null || authority.isEmpty()) {
throw new IllegalArgumentException("HTTP/2 URI requires an authority");
}
HpackEncoder.writeLiteralWithNameIndex(
outgoing, 1, authority.getBytes(StandardCharsets.US_ASCII), false);
}
private void writePath(URI uri) {
String path = uri.getRawPath();
if (path == null || path.isEmpty()) path = "/";
if (uri.getRawQuery() != null) path += "?" + uri.getRawQuery();
if ("/".equals(path)) HpackEncoder.writeIndexed(outgoing, 4);
else if ("/index.html".equals(path)) HpackEncoder.writeIndexed(outgoing, 5);
else {
HpackEncoder.writeLiteralWithNameIndex(
outgoing, 4, path.getBytes(StandardCharsets.US_ASCII), false);
}
}
@Override
public synchronized void close() {
if (closed) return;
closed = true;
writer.close();
try {
socket.close();
} catch (IOException ignored) {
// Closing a broken pooled connection is best-effort.
}
}
private static Socket connect(Origin origin, SSLContext sslContext) throws IOException {
if (!origin.secure) {
Socket socket = new Socket();
socket.connect(new InetSocketAddress(origin.host, origin.port), CONNECT_TIMEOUT_MS);
return socket;
}
SSLContext context;
try {
context = sslContext == null ? SSLContext.getDefault() : sslContext;
} catch (Exception failure) {
throw new IOException("cannot initialize TLS context", failure);
}
SSLSocket socket =
(SSLSocket) context.getSocketFactory().createSocket(origin.host, origin.port);
SSLParameters parameters = socket.getSSLParameters();
parameters.setApplicationProtocols(new String[] {"h2"});
parameters.setEndpointIdentificationAlgorithm("HTTPS");
socket.setSSLParameters(parameters);
socket.startHandshake();
if (!"h2".equals(socket.getApplicationProtocol())) {
socket.close();
throw new IOException("upstream did not negotiate HTTP/2 through ALPN");
}
return socket;
}
}
private static final class Exchange {
private final int streamId;
private final MutableHeaderMap headers = new MutableHeaderMap();
private final MutableHeaderMap trailers = new MutableHeaderMap();
private final ByteArrayOutputStream body = new ByteArrayOutputStream();
private int statusCode;
private boolean complete;
private Exchange(int streamId) {
this.streamId = streamId;
}
private Http2ClientResponse response() {
return new Http2ClientResponse(statusCode, headers, body.toByteArray(), trailers);
}
}
private static final class ResponseHeaderSink
implements dev.relism.flash.http2.hpack.HeaderSink {
private final Exchange exchange;
private final boolean trailers;
private boolean regular;
private boolean status;
private ResponseHeaderSink(Exchange exchange, boolean trailers) {
this.exchange = exchange;
this.trailers = trailers;
}
@Override
public void accept(ByteView name, ByteView value, boolean neverIndexed) {
if (name.length() != 0 && name.byteAt(0) == ':') {
if (trailers || regular || status || !equalsAscii(name, ":status")) {
throw Http2Exception.PROTOCOL_ERROR;
}
exchange.statusCode = parseStatus(value);
status = true;
return;
}
regular = true;
MutableHeaderMap target = trailers ? exchange.trailers : exchange.headers;
byte[] nameBytes = copy(name);
byte[] valueBytes = copy(value);
target.add(nameBytes, 0, nameBytes.length, valueBytes, 0, valueBytes.length);
}
private void validate() throws IOException {
if (!trailers && !status) throw new IOException("HTTP/2 response omitted :status");
}
private static int parseStatus(ByteView value) {
if (value.length() != 3) throw Http2Exception.PROTOCOL_ERROR;
int code = 0;
for (int i = 0; i < 3; i++) {
int digit = (value.byteAt(i) & 0xff) - '0';
if (digit < 0 || digit > 9) throw Http2Exception.PROTOCOL_ERROR;
code = code * 10 + digit;
}
return code;
}
}
private static final class BufferIntent implements WriteIntent {
private byte[] bytes;
private int length;
private WriteIntent next;
private void reset(byte[] bytes, int length) {
this.bytes = bytes;
this.length = length;
this.next = null;
}
@Override public byte[] buffer() { return bytes; }
@Override public int offset() { return 0; }
@Override public int length() { return length; }
@Override public WriteIntent mpscNext() { return next; }
@Override public void setMpscNext(WriteIntent next) { this.next = next; }
}
private record Origin(String scheme, String host, int port, boolean secure) {
private static Origin from(URI uri) {
String scheme = uri.getScheme();
boolean secure;
if ("https".equalsIgnoreCase(scheme)) secure = true;
else if ("http".equalsIgnoreCase(scheme)) secure = false;
else throw new IllegalArgumentException("HTTP/2 URI scheme must be http or https");
if (uri.getHost() == null) throw new IllegalArgumentException("HTTP/2 URI requires a host");
int port = uri.getPort() >= 0 ? uri.getPort() : secure ? 443 : 80;
return new Origin(scheme.toLowerCase(), uri.getHost(), port, secure);
}
}
private static final class OpenFailure extends RuntimeException {
private final IOException io;
private OpenFailure(IOException io) {
super(io);
this.io = io;
}
}
private static boolean equalsAscii(ByteView bytes, String value) {
if (bytes.length() != value.length()) return false;
for (int i = 0; i < bytes.length(); i++) {
int left = bytes.byteAt(i) & 0xff;
int right = value.charAt(i);
if (left >= 'A' && left <= 'Z') left += 'a' - 'A';
if (right >= 'A' && right <= 'Z') right += 'a' - 'A';
if (left != right) return false;
}
return true;
}
private static byte[] copy(ByteView view) {
byte[] result = new byte[view.length()];
for (int i = 0; i < result.length; i++) result[i] = view.byteAt(i);
return result;
}
private static int addWindow(int current, int increment) throws IOException {
long next = (long) current + increment;
if (next > Integer.MAX_VALUE) throw new IOException("HTTP/2 flow-control window overflow");
return (int) next;
}
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);
}
}
@@ -0,0 +1,7 @@
package dev.relism.flash.http2.client;
import dev.relism.flash.models.HeaderView;
/** Complete response returned by Flash's proxy-oriented HTTP/2 client. */
public record Http2ClientResponse(
int statusCode, HeaderView headers, byte[] body, HeaderView trailers) {}
@@ -1,6 +1,7 @@
package dev.relism.flash.http2.hpack;
import dev.relism.flash.bytes.ByteWriter;
import dev.relism.fpr.core.ByteView;
/**
* Stateless HPACK encoder for response header blocks. It uses the RFC 7541 static table and literal
@@ -25,6 +26,19 @@ public final class HpackEncoder {
writeLiteral(out, name, 0, name.length, value, 0, value.length, false);
}
/** Writes a non-indexed literal directly from protocol-neutral byte views. */
public static void writeLiteral(ByteWriter out, ByteView name, ByteView value) {
HpackIntegers.encode(out, 0, 4, 0);
HpackIntegers.encode(out, 0, 7, name.length());
for (int i = 0; i < name.length(); i++) {
int octet = name.byteAt(i) & 0xff;
if (octet >= 'A' && octet <= 'Z') octet += 'a' - 'A';
out.writeByte((byte) octet);
}
HpackIntegers.encode(out, 0, 7, value.length());
for (int i = 0; i < value.length(); i++) out.writeByte(value.byteAt(i));
}
public static void writeLiteral(
ByteWriter out,
byte[] name,
@@ -134,17 +134,13 @@ public final class ConnectionRunner {
}
}
/**
* Decides h1 vs h2 for this connection, applying {@link FlashConfiguration#isHttp2Enabled()} to
* the plaintext (h2c) path — see {@link ProtocolNegotiator}'s Javadoc for why the flag is applied
* here rather than inside the negotiator itself.
*/
/** Decides h1 vs h2 while keeping the TLS and cleartext rollout gates independent. */
private NegotiatedProtocol negotiateProtocol(Socket socket, BufferedByteSource in)
throws IOException {
if (socket instanceof SSLSocket) {
return ProtocolNegotiator.negotiate(socket, in); // ALPN — already resolved, no I/O
}
if (!configuration.isHttp2Enabled()) {
if (!configuration.isHttp2CleartextEnabled()) {
return NegotiatedProtocol.HTTP_1_1; // skip the h2c peek entirely when disabled
}
in.setDeadline(System.nanoTime() + configuration.getIdleKeepAliveTimeoutMs() * 1_000_000L);
@@ -1,63 +1,35 @@
package dev.relism.flash.transport;
import javax.net.ssl.SSLSocket;
import java.io.IOException;
import java.net.Socket;
import java.nio.charset.StandardCharsets;
import java.util.Arrays;
import javax.net.ssl.SSLSocket;
/**
* Decides, once per connection and before any request is parsed, whether the connection speaks
* immediately after ALPN/preface detection").
*
* <p>Two independent signals, in order:
* <ol>
* <li><b>ALPN</b> (TLS connections). If the socket is an {@link SSLSocket} and the TLS
* handshake already resolved {@code "h2"} as the application protocol, this connection is
* {@link NegotiatedProtocol#HTTP_2}. Anything else negotiated — {@code "http/1.1"}, no
* protocol at all (a peer that doesn't speak ALPN), or an empty string — is
* {@link NegotiatedProtocol#HTTP_1_1}. This costs nothing beyond a field read: ALPN is
* resolved during the handshake, which must already have completed (see
* {@code TlsConfig}'s Javadoc on why {@code startHandshake()} must be called explicitly
* <li><b>h2c prior knowledge</b> (plaintext connections, RFC 9113 §3.4). The first 24 bytes of
* the connection are compared, without being consumed, against the client connection
* preface {@code "PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"}. A match is
* {@link NegotiatedProtocol#HTTP_2}; anything else — including a partial match followed by
* EOF, or a preface look-alike that diverges partway through — is
* {@link NegotiatedProtocol#HTTP_1_1}. This is why {@link BufferedByteSource#peek} exists:
* the bytes must remain available for {@code RequestParser} if they turn out not to be an
* h2 preface after all.</li>
* </ol>
*
* <p>This method reports the protocol accurately and unconditionally — it does not consult
* {@code FlashConfiguration.http2Enabled}. Gating whether an {@link NegotiatedProtocol#HTTP_2}
* {@code Http2Connection} yet) and whether the h2c peek is even attempted for plaintext
* connections are both the caller's responsibility, so that this class stays a pure,
* directly-testable detector (see {@code ProtocolNegotiatorTest}).
*/
/** Detects HTTP/1.1 or HTTP/2 once, before the connection parser is selected. */
public final class ProtocolNegotiator {
private static final byte[] H2C_PREFACE =
"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n".getBytes(StandardCharsets.US_ASCII);
/**
* reconstructed per connection.
*/
private static final byte[] H2C_PREFACE =
"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n".getBytes(StandardCharsets.US_ASCII);
private ProtocolNegotiator() {}
private ProtocolNegotiator() {
/**
* Uses the completed TLS ALPN result for secure sockets and a non-consuming prior-knowledge
* preface probe for plaintext sockets. Configuration gates remain the caller's responsibility,
* which keeps detection deterministic and independently testable.
*/
public static NegotiatedProtocol negotiate(Socket socket, BufferedByteSource source)
throws IOException {
if (socket instanceof SSLSocket ssl) {
return "h2".equals(ssl.getApplicationProtocol())
? NegotiatedProtocol.HTTP_2
: NegotiatedProtocol.HTTP_1_1;
}
public static NegotiatedProtocol negotiate(Socket socket, BufferedByteSource source) throws IOException {
if (socket instanceof SSLSocket ssl) {
String applicationProtocol = ssl.getApplicationProtocol();
return "h2".equals(applicationProtocol) ? NegotiatedProtocol.HTTP_2 : NegotiatedProtocol.HTTP_1_1;
}
byte[] probe = new byte[H2C_PREFACE.length];
int n = source.peek(probe, 0, probe.length);
if (n == H2C_PREFACE.length && Arrays.equals(probe, H2C_PREFACE)) {
return NegotiatedProtocol.HTTP_2;
}
return NegotiatedProtocol.HTTP_1_1;
}
byte[] probe = new byte[H2C_PREFACE.length];
int read = source.peek(probe, 0, probe.length);
return read == H2C_PREFACE.length && Arrays.equals(probe, H2C_PREFACE)
? NegotiatedProtocol.HTTP_2
: NegotiatedProtocol.HTTP_1_1;
}
}
@@ -0,0 +1,55 @@
package dev.relism.flash.http;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import dev.relism.flash.bytes.PooledSlice;
import dev.relism.flash.http.HopByHopHeaders.Protocol;
import dev.relism.flash.models.MutableHeaderMap;
import java.nio.charset.StandardCharsets;
import org.junit.jupiter.api.Test;
class HopByHopHeaderTest {
@Test
void sharedPolicyCoversAllFourProtocolConversions() {
for (Protocol sourceProtocol : Protocol.values()) {
for (Protocol targetProtocol : Protocol.values()) {
MutableHeaderMap source = new MutableHeaderMap();
add(source, "connection", "x-private, keep-alive");
add(source, "x-private", "secret");
add(source, "upgrade", "websocket");
add(source, "te", "trailers");
add(source, "x-end-to-end", "yes");
assertFalse(forward(source, "connection", "x-private", sourceProtocol, targetProtocol));
assertFalse(forward(source, "x-private", "secret", sourceProtocol, targetProtocol));
assertFalse(forward(source, "upgrade", "websocket", sourceProtocol, targetProtocol));
assertTrue(forward(source, "x-end-to-end", "yes", sourceProtocol, targetProtocol));
assertTrue(forward(source, "te", "trailers", sourceProtocol, Protocol.HTTP_2));
assertFalse(forward(source, "te", "trailers", sourceProtocol, Protocol.HTTP_1_1));
}
}
}
private static boolean forward(
MutableHeaderMap source,
String name,
String value,
Protocol sourceProtocol,
Protocol targetProtocol) {
byte[] nameBytes = name.getBytes(StandardCharsets.US_ASCII);
byte[] valueBytes = value.getBytes(StandardCharsets.US_ASCII);
PooledSlice nameView = new PooledSlice();
PooledSlice valueView = new PooledSlice();
nameView.reset(nameBytes, 0, nameBytes.length);
valueView.reset(valueBytes, 0, valueBytes.length);
return HopByHopHeaders.shouldForward(
source, nameView, valueView, sourceProtocol, targetProtocol);
}
private static void add(MutableHeaderMap headers, String name, String value) {
byte[] nameBytes = name.getBytes(StandardCharsets.US_ASCII);
byte[] valueBytes = value.getBytes(StandardCharsets.US_ASCII);
headers.add(nameBytes, 0, nameBytes.length, valueBytes, 0, valueBytes.length);
}
}
@@ -29,8 +29,13 @@ class GrpcInteropTest {
@Test
void grpcurlCompletesUnaryStreamingAndErrorCalls(@TempDir Path directory) throws Exception {
int port = freePort();
app = FlashApp.create(FlashConfiguration.builder()
.host("127.0.0.1").port(port).http2Enabled(true).build());
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.post("/flash.test.Echo/Unary", (request, response) ->
response.type("application/grpc")
.body(request.body().bytes())
@@ -0,0 +1,71 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertEquals;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
import dev.relism.flash.http2.client.Http2Client;
import dev.relism.flash.http2.client.Http2ClientResponse;
import java.io.InputStream;
import java.net.ServerSocket;
import java.net.Socket;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
class H2cPriorKnowledgeTest {
private FlashApp app;
@AfterEach
void stop() {
if (app != null) app.stop().join();
}
@Test
void priorKnowledgeRequiresItsIndependentOptIn() throws Exception {
int disabledPort = freePort();
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(disabledPort)
.http2Enabled(true)
.build());
app.get("/", (request, response) -> "wrong protocol");
app.start();
try (Socket socket = new Socket("127.0.0.1", disabledPort)) {
socket.setSoTimeout(2_000);
socket.getOutputStream().write(Http2Preface.clientPreface());
byte[] prefix = socket.getInputStream().readNBytes(5);
assertArrayEquals("HTTP/".getBytes(StandardCharsets.US_ASCII), prefix);
}
app.stop().join();
int enabledPort = freePort();
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(enabledPort)
.http2CleartextEnabled(true)
.build());
app.get("/", (request, response) -> "h2c");
app.start();
try (Http2Client client = new Http2Client()) {
Http2ClientResponse response =
client.get(URI.create("http://127.0.0.1:" + enabledPort + "/"));
assertEquals(200, response.statusCode());
assertEquals("h2c", new String(response.body(), StandardCharsets.UTF_8));
}
}
private static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -184,8 +184,14 @@ class Http2AbuseTest {
@Test
void idleOpenStreamIsCancelledWithinConfiguredDeadline() throws Exception {
int port = freePort();
FlashApp app = FlashApp.create(FlashConfiguration.builder()
.host("127.0.0.1").port(port).http2Enabled(true).h2StreamIdleTimeoutMs(20).build());
FlashApp app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.h2StreamIdleTimeoutMs(20)
.build());
app.post("/idle", (request, response) -> request.body().bytes());
app.start();
ByteWriter headers = new ByteWriter(64);
@@ -0,0 +1,18 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import org.junit.jupiter.api.Test;
class Http2AuthorityTest {
@Test
void matchesExactIpPortAndSingleLabelWildcardAuthorities() {
assertTrue(Http2Authority.matches("api.example.com:443", "api.example.com"));
assertTrue(Http2Authority.matches("127.0.0.1:8443", "127.0.0.1"));
assertTrue(Http2Authority.matches("one.example.com", "*.example.com"));
assertFalse(Http2Authority.matches("example.com", "*.example.com"));
assertFalse(Http2Authority.matches("two.one.example.com", "*.example.com"));
assertFalse(Http2Authority.matches("other.example.net", "*.example.com"));
}
}
@@ -29,8 +29,13 @@ class Http2ConnectTest {
@Test
void connectTunnelCanExchangeDataBeforeEitherSideCloses() throws Exception {
int port = freePort();
app = FlashApp.create(FlashConfiguration.builder()
.host("127.0.0.1").port(port).http2Enabled(true).build());
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.connect("tunnel", (request, response) ->
response.type(ContentType.NONE).streaming(output -> {
byte[] bytes = new byte[16];
@@ -253,7 +253,11 @@ class Http2ConnectionIntegrationTest {
int port = freePort();
app =
FlashApp.create(
FlashConfiguration.builder().port(port).host("127.0.0.1").http2Enabled(true).build());
FlashConfiguration.builder()
.port(port)
.host("127.0.0.1")
.http2CleartextEnabled(true)
.build());
app.get("/api/ping", (request, response) -> "pong");
app.start();
@@ -314,7 +318,11 @@ class Http2ConnectionIntegrationTest {
AtomicInteger calls = new AtomicInteger();
app =
FlashApp.create(
FlashConfiguration.builder().port(port).host("127.0.0.1").http2Enabled(true).build());
FlashConfiguration.builder()
.port(port)
.host("127.0.0.1")
.http2CleartextEnabled(true)
.build());
app.get(
"/queued",
(request, response) -> {
@@ -380,7 +388,11 @@ class Http2ConnectionIntegrationTest {
AtomicBoolean handlerEntered = new AtomicBoolean();
app =
FlashApp.create(
FlashConfiguration.builder().port(port).host("127.0.0.1").http2Enabled(true).build());
FlashConfiguration.builder()
.port(port)
.host("127.0.0.1")
.http2CleartextEnabled(true)
.build());
app.get(
"/",
(request, response) -> {
@@ -436,7 +448,7 @@ class Http2ConnectionIntegrationTest {
FlashConfiguration.builder()
.port(port)
.host("127.0.0.1")
.http2Enabled(true)
.http2CleartextEnabled(true)
.shutdownDrainTimeoutMs(5_000)
.build());
app.start();
@@ -476,7 +488,11 @@ class Http2ConnectionIntegrationTest {
int port = freePort();
app =
FlashApp.create(
FlashConfiguration.builder().port(port).host("127.0.0.1").http2Enabled(true).build());
FlashConfiguration.builder()
.port(port)
.host("127.0.0.1")
.http2CleartextEnabled(true)
.build());
app.start();
try (Socket first = new Socket("127.0.0.1", port)) {
@@ -0,0 +1,122 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertEquals;
import dev.relism.flash.bytes.ByteWriter;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
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.tls.TestKeystores;
import dev.relism.flash.tls.TlsConfig;
import java.io.InputStream;
import java.net.ServerSocket;
import java.nio.file.Path;
import javax.net.ssl.SSLParameters;
import javax.net.ssl.SSLSocket;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.io.TempDir;
class Http2MisdirectedRequestTest {
private FlashApp app;
@AfterEach
void stop() {
if (app != null) app.stop().join();
}
@Test
void authorityOutsideSelectedCertificateReceives421(@TempDir Path directory) throws Exception {
int port = freePort();
Path keystore =
TestKeystores.build(
directory,
"misdirected.p12",
"changeit",
TestKeystores.Entry.of("server", "localhost", "localhost"));
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.tls(TlsConfig.keystore(keystore, "changeit"))
.http2Enabled(true)
.build());
app.get("/", (request, response) -> "must not run");
app.start();
try (SSLSocket socket =
(SSLSocket)
TestKeystores.trustAllClientContext()
.getSocketFactory()
.createSocket("localhost", port)) {
SSLParameters parameters = socket.getSSLParameters();
parameters.setApplicationProtocols(new String[] {"h2"});
socket.setSSLParameters(parameters);
socket.startHandshake();
socket.getOutputStream().write(request("other.example"));
assertEquals(421, readStatus(socket.getInputStream()));
}
}
private static byte[] request(String authority) {
ByteWriter bytes = new ByteWriter(128);
bytes.writeBytes(Http2Preface.clientPreface());
FrameWriteBuffer frames = new FrameWriteBuffer(bytes);
frames.beginFrame(FrameType.SETTINGS, 0, 0);
frames.endFrame();
frames.beginFrame(
FrameType.HEADERS, FrameFlags.END_HEADERS | FrameFlags.END_STREAM, 1);
HpackEncoder.writeIndexed(bytes, 2);
HpackEncoder.writeIndexed(bytes, 7);
HpackEncoder.writeLiteralWithNameIndex(
bytes, 1, authority.getBytes(java.nio.charset.StandardCharsets.US_ASCII), false);
HpackEncoder.writeIndexed(bytes, 4);
frames.endFrame();
byte[] result = new byte[bytes.length()];
System.arraycopy(bytes.array(), 0, result, 0, result.length);
return result;
}
private static int readStatus(InputStream input) throws Exception {
HpackDecoder decoder = new HpackDecoder();
byte[] header = new byte[9];
while (true) {
input.readNBytes(header, 0, header.length);
int length = ((header[0] & 0xff) << 16) | ((header[1] & 0xff) << 8) | (header[2] & 0xff);
int type = header[3] & 0xff;
int streamId =
((header[5] & 0x7f) << 24)
| ((header[6] & 0xff) << 16)
| ((header[7] & 0xff) << 8)
| (header[8] & 0xff);
byte[] payload = input.readNBytes(length);
if (type != FrameType.HEADERS.code() || streamId != 1) continue;
int[] status = {0};
decoder.decode(
payload,
0,
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 static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -28,8 +28,13 @@ class Http2TrailersTest {
@Test
void requestTrailersReachHandlerAfterBodyEof() throws Exception {
int port = freePort();
app = FlashApp.create(FlashConfiguration.builder()
.host("127.0.0.1").port(port).http2Enabled(true).build());
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.post("/trailers", (request, response) -> {
assertEquals("abc", new String(request.body().bytes(), StandardCharsets.US_ASCII));
return request.trailers().first("grpc-status");
@@ -95,8 +100,13 @@ class Http2TrailersTest {
private int startBlockingRoute() throws Exception {
int port = freePort();
app = FlashApp.create(FlashConfiguration.builder()
.host("127.0.0.1").port(port).http2Enabled(true).build());
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.post("/trailers", (request, response) -> request.body().bytes());
app.start();
return port;
@@ -0,0 +1,131 @@
package dev.relism.flash.http2;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.http.proxy.HttpProxy;
import dev.relism.flash.http2.client.Http2Client;
import dev.relism.flash.http2.client.Http2ClientResponse;
import dev.relism.flash.models.MutableHeaderMap;
import java.io.ByteArrayOutputStream;
import java.net.ServerSocket;
import java.net.Socket;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
class ProxyTrailerRelayTest {
private FlashApp upstream;
private FlashApp proxy;
private Http2Client proxyUpstream;
@AfterEach
void stop() {
if (proxyUpstream != null) proxyUpstream.close();
if (proxy != null) proxy.stop().join();
if (upstream != null) upstream.stop().join();
}
@Test
void requestAndResponseTrailersSurviveH2AndH1DownstreamProxyHops() throws Exception {
int upstreamPort = freePort();
upstream =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(upstreamPort)
.http2CleartextEnabled(true)
.build());
upstream.post(
"/relay",
(request, response) ->
response
.header("x-query", request.query("mode"))
.header("x-private-seen", String.valueOf(request.header("x-private") != null))
.body(request.body().bytes())
.trailer("x-relayed-trailer", request.trailers().first("x-request-trailer")));
upstream.start();
int proxyPort = freePort();
proxyUpstream = new Http2Client();
proxy =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(proxyPort)
.http2CleartextEnabled(true)
.build());
proxy.post(
"/relay",
HttpProxy.toHttp2(URI.create("http://127.0.0.1:" + upstreamPort), proxyUpstream));
proxy.start();
MutableHeaderMap h2Headers = fields("connection", "x-private");
add(h2Headers, "x-private", "must-not-cross");
MutableHeaderMap h2Trailers = fields("x-request-trailer", "from-h2");
try (Http2Client downstream = new Http2Client()) {
Http2ClientResponse response =
downstream.exchange(
URI.create("http://127.0.0.1:" + proxyPort + "/relay?mode=h2"),
HttpMethod.POST,
h2Headers,
"hello-h2".getBytes(StandardCharsets.UTF_8),
h2Trailers);
assertEquals("hello-h2", new String(response.body(), StandardCharsets.UTF_8));
assertEquals("h2", response.headers().first("x-query"));
assertEquals("false", response.headers().first("x-private-seen"));
assertEquals("from-h2", response.trailers().first("x-relayed-trailer"));
}
String h1 = h1Exchange(proxyPort);
assertTrue(h1.contains("hello-h1"), h1);
assertTrue(h1.toLowerCase().contains("x-query: h1"), h1);
assertTrue(h1.toLowerCase().contains("x-private-seen: false"), h1);
assertTrue(h1.toLowerCase().contains("x-relayed-trailer: from-h1"), h1);
assertFalse(h1.contains("must-not-cross"), h1);
}
private static String h1Exchange(int port) throws Exception {
try (Socket socket = new Socket("127.0.0.1", port)) {
socket.setSoTimeout(2_000);
socket
.getOutputStream()
.write(
("POST /relay?mode=h1 HTTP/1.1\r\n"
+ "Host: 127.0.0.1\r\n"
+ "Connection: x-private, close\r\n"
+ "X-Private: must-not-cross\r\n"
+ "Transfer-Encoding: chunked\r\n"
+ "Trailer: x-request-trailer\r\n\r\n"
+ "8\r\nhello-h1\r\n"
+ "0\r\nX-Request-Trailer: from-h1\r\n\r\n")
.getBytes(StandardCharsets.US_ASCII));
ByteArrayOutputStream bytes = new ByteArrayOutputStream();
socket.getInputStream().transferTo(bytes);
return bytes.toString(StandardCharsets.UTF_8);
}
}
private static MutableHeaderMap fields(String name, String value) {
MutableHeaderMap headers = new MutableHeaderMap();
add(headers, name, value);
return headers;
}
private static void add(MutableHeaderMap headers, String name, String value) {
byte[] nameBytes = name.getBytes(StandardCharsets.US_ASCII);
byte[] valueBytes = value.getBytes(StandardCharsets.UTF_8);
headers.add(nameBytes, 0, nameBytes.length, valueBytes, 0, valueBytes.length);
}
private static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -0,0 +1,117 @@
package dev.relism.flash.http2.client;
import static org.junit.jupiter.api.Assertions.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertEquals;
import dev.relism.flash.extension.FlashApp;
import dev.relism.flash.extension.FlashConfiguration;
import dev.relism.flash.http.HttpMethod;
import dev.relism.flash.models.MutableHeaderMap;
import dev.relism.flash.tls.TestKeystores;
import dev.relism.flash.tls.TlsConfig;
import java.net.ServerSocket;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.nio.file.Path;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.io.TempDir;
class Http2ClientTest {
private FlashApp app;
@AfterEach
void stop() {
if (app != null) app.stop().join();
}
@Test
void reusesOriginConnectionAndExchangesFlowControlledBodiesAndTrailers() throws Exception {
int port = freePort();
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.http2CleartextEnabled(true)
.build());
app.post(
"/relay",
(request, response) -> {
byte[] body = request.body().bytes();
String checksum = request.trailers().first("x-request-checksum");
return response
.header("x-upstream", request.header("x-forwarded-test"))
.body(body)
.trailer("x-response-checksum", checksum);
});
app.start();
byte[] body = new byte[2 * 1024 * 1024 + 31];
for (int i = 0; i < body.length; i++) body[i] = (byte) (i * 29);
MutableHeaderMap requestHeaders = fields("x-forwarded-test", "yes");
MutableHeaderMap requestTrailers = fields("x-request-checksum", "valid");
try (Http2Client client = new Http2Client()) {
URI uri = URI.create("http://127.0.0.1:" + port + "/relay");
Http2ClientResponse first =
client.exchange(uri, HttpMethod.POST, requestHeaders, body, requestTrailers);
Http2ClientResponse second =
client.exchange(
uri,
HttpMethod.POST,
requestHeaders,
"again".getBytes(StandardCharsets.UTF_8),
requestTrailers);
assertEquals(200, first.statusCode());
assertEquals("yes", first.headers().first("x-upstream"));
assertArrayEquals(body, first.body());
assertEquals("valid", first.trailers().first("x-response-checksum"));
assertArrayEquals("again".getBytes(StandardCharsets.UTF_8), second.body());
assertEquals(1, client.pooledConnectionCount());
}
}
@Test
void negotiatesTlsAlpnAndVerifiesTheUpstreamHostname(@TempDir Path directory) throws Exception {
int port = freePort();
Path keystore =
TestKeystores.build(
directory,
"http2-client.p12",
"changeit",
TestKeystores.Entry.of("server", "localhost", "localhost"));
app =
FlashApp.create(
FlashConfiguration.builder()
.host("127.0.0.1")
.port(port)
.tls(TlsConfig.keystore(keystore, "changeit"))
.http2Enabled(true)
.build());
app.get("/secure", (request, response) -> "tls-h2");
app.start();
try (Http2Client client = new Http2Client(TestKeystores.trustAllClientContext())) {
Http2ClientResponse response =
client.get(URI.create("https://localhost:" + port + "/secure"));
assertEquals(200, response.statusCode());
assertEquals("tls-h2", new String(response.body(), StandardCharsets.UTF_8));
}
}
private static MutableHeaderMap fields(String name, String value) {
MutableHeaderMap headers = new MutableHeaderMap();
byte[] nameBytes = name.getBytes(StandardCharsets.US_ASCII);
byte[] valueBytes = value.getBytes(StandardCharsets.UTF_8);
headers.add(nameBytes, 0, nameBytes.length, valueBytes, 0, valueBytes.length);
return headers;
}
private static int freePort() throws Exception {
try (ServerSocket socket = new ServerSocket(0)) {
return socket.getLocalPort();
}
}
}
@@ -21,7 +21,7 @@ import static org.junit.jupiter.api.Assertions.*;
/**
* {@link ProtocolNegotiator#negotiate} is a pure, directly-testable detector (see its Javadoc
* for why it does not itself consult {@code FlashConfiguration.http2Enabled}) — every case here
* for why it does not itself consult {@code FlashConfiguration}) — every case here
* calls it directly rather than through {@code Http1Connection}/{@code ConnectionRunner}.
*/
class ProtocolNegotiatorTest {