diff --git a/README.md b/README.md index b0bb1c3..33a6379 100644 --- a/README.md +++ b/README.md @@ -137,6 +137,15 @@ public final class PublishedTunnel { `forwardTo()` remains available for existing services that already bind a local TCP port. +The control channel uses a negotiated heartbeat deadline. A missed deadline or +an unexpected control-transport EOF closes tunnel admission immediately, while +streams already returned by `accept()` and active `forwardTo()` relays keep +their independent payload sockets and drain normally. This prevents a +transient or asymmetric control-path outage from interrupting an established +application session. Explicit tunnel/control closure and protocol revocation +remain hard lifecycle operations enforced by the engine. Treat stream EOF or +I/O failure—not `control.done()`—as the payload-lifecycle signal. + ## Private dial ```java diff --git a/src/main/java/io/rstream/BytestreamTunnel.java b/src/main/java/io/rstream/BytestreamTunnel.java index 7626681..5d2d8e7 100644 --- a/src/main/java/io/rstream/BytestreamTunnel.java +++ b/src/main/java/io/rstream/BytestreamTunnel.java @@ -5,21 +5,35 @@ import java.io.OutputStream; import java.net.Socket; import java.time.Duration; +import java.util.ArrayDeque; +import java.util.ArrayList; +import java.util.Collections; +import java.util.HashSet; import java.util.Objects; -import java.util.concurrent.BlockingQueue; +import java.util.Set; +import java.util.WeakHashMap; +import java.util.concurrent.CancellationException; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.ExecutionException; import java.util.concurrent.ExecutorService; -import java.util.concurrent.LinkedBlockingQueue; -import java.util.concurrent.TimeUnit; +import java.util.concurrent.FutureTask; +import java.util.concurrent.RejectedExecutionException; +import java.util.concurrent.locks.Condition; +import java.util.concurrent.locks.ReentrantLock; /** A bytestream tunnel opened on the rstream engine. */ public final class BytestreamTunnel implements AutoCloseable { - private static final Object CLOSED = new Object(); private final ControlChannel control; private final TunnelProperties properties; - private final BlockingQueue streams = new LinkedBlockingQueue<>(); + private final StreamQueue streams = new StreamQueue(); private final ExecutorService executor; + private final Object lifecycleLock = new Object(); + private final Set> forwarders = new HashSet<>(); + // Application-owned accepted streams must not be retained solely by lifecycle tracking. + private final Set payloads = Collections.newSetFromMap(new WeakHashMap<>()); + private volatile boolean hardClosed; private volatile boolean closed; + private RuntimeException hardCloseError; BytestreamTunnel(ControlChannel control, TunnelProperties properties, ExecutorService executor) { if (properties.id() == null || properties.id().isBlank()) { @@ -47,15 +61,13 @@ public String forwardingAddress() { } public RstreamStream accept() throws InterruptedException { - return accepted(streams.take()); + return finishAcceptance(streams.take(null)); } public RstreamStream accept(Duration timeout) throws InterruptedException { Objects.requireNonNull(timeout, "timeout"); if (timeout.isNegative()) throw new IllegalArgumentException("timeout must not be negative"); - var item = streams.poll(timeout.toMillis(), TimeUnit.MILLISECONDS); - if (item == null) return null; - return accepted(item); + return finishAcceptance(streams.take(timeout)); } public CompletableFuture acceptAsync() { @@ -83,7 +95,7 @@ public CompletableFuture forwardTo(String host, int port) { while (!closed) { try { var stream = accept(); - executor.submit(() -> pipeToLocal(stream, host, port)); + startForwarder(stream, host, port); } catch (InterruptedException error) { Thread.currentThread().interrupt(); return; @@ -105,25 +117,58 @@ public CompletableFuture serveHttp(RstreamHttpHandler handler, RstreamHttp } boolean deliver(RstreamStream stream) { - if (closed) { - stream.closeQuietly(); - return false; + synchronized (lifecycleLock) { + if (closed) { + stream.closeQuietly(); + return false; + } + payloads.add(stream); } - streams.offer(stream); - return true; + return streams.offer(stream); } void onClose(Throwable error) { - if (closed) return; - closed = true; - if (error != null) streams.offer(error); - streams.offer(CLOSED); + onClose(error, false); + } + + void onClose(Throwable error, boolean preserveForwarders) { + var closeError = + error instanceof RuntimeException runtimeError + ? runtimeError + : new RstreamException("Tunnel closed.", "ERR_RSTREAM_TUNNEL_CLOSED", error); + var firstClose = false; + Set> canceled = Set.of(); + Set closedPayloads = Set.of(); + synchronized (lifecycleLock) { + if (!closed) { + closed = true; + firstClose = true; + } + if (!preserveForwarders && !hardClosed) { + hardClosed = true; + hardCloseError = closeError; + canceled = Set.copyOf(forwarders); + closedPayloads = Set.copyOf(payloads); + payloads.clear(); + } + } + if (firstClose) streams.close(closeError); + closedPayloads.forEach(RstreamStream::closeQuietly); + canceled.forEach(task -> task.cancel(true)); } @Override public void close() { - if (closed) return; - control.closeTunnel(id()); + if (closed) { + onClose(null, false); + return; + } + try { + control.closeTunnel(id()); + } catch (RuntimeException error) { + if (closed) onClose(error, false); + throw error; + } } public CompletableFuture closeAsync() { @@ -149,17 +194,6 @@ static String formatForwardingAddress(TunnelProperties properties) { "Invalid tunnel properties: no host, name, or ID.", "ERR_RSTREAM_INVALID_TUNNEL"); } - private static RstreamStream accepted(Object item) { - if (item == CLOSED) { - throw new RstreamException("Tunnel closed.", "ERR_RSTREAM_TUNNEL_CLOSED"); - } - if (item instanceof RuntimeException error) throw error; - if (item instanceof Throwable error) { - throw new RstreamException("Tunnel closed.", "ERR_RSTREAM_TUNNEL_CLOSED", error); - } - return (RstreamStream) item; - } - private static String publishedHost(TunnelProperties properties) { if (properties.hostname() != null && !properties.hostname().isBlank()) { var port = properties.port() == null ? 443 : properties.port(); @@ -208,6 +242,57 @@ private static void pipeToLocal(RstreamStream stream, String host, int port) { } } + private void startForwarder(RstreamStream stream, String host, int port) { + var task = + new FutureTask(() -> pipeToLocal(stream, host, port), null) { + @Override + protected void done() { + forwarderDone(this); + } + }; + synchronized (lifecycleLock) { + if (hardClosed) { + stream.closeQuietly(); + return; + } + forwarders.add(task); + } + try { + executor.execute(task); + } catch (RejectedExecutionException error) { + synchronized (lifecycleLock) { + forwarders.remove(task); + } + stream.closeQuietly(); + throw new RstreamException( + "Failed to schedule tunnel forwarding.", "ERR_RSTREAM_FORWARD", error); + } + } + + private RstreamStream finishAcceptance(RstreamStream stream) { + if (stream == null) return null; + RuntimeException closeError; + synchronized (lifecycleLock) { + closeError = hardClosed ? hardCloseError : null; + } + if (closeError == null) return stream; + stream.closeQuietly(); + throw closeError; + } + + private void forwarderDone(FutureTask task) { + synchronized (lifecycleLock) { + forwarders.remove(task); + } + try { + task.get(); + } catch (CancellationException ignored) { + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + } catch (ExecutionException ignored) { + } + } + private static void copy(InputStream input, OutputStream output) { try { input.transferTo(output); @@ -222,4 +307,66 @@ private static void shutdownOutput(Socket socket) { } catch (IOException ignored) { } } + + private static final class StreamQueue { + private final ReentrantLock lock = new ReentrantLock(); + private final Condition available = lock.newCondition(); + private final ArrayDeque queued = new ArrayDeque<>(); + private RuntimeException closeError; + + RstreamStream take(Duration timeout) throws InterruptedException { + var remaining = timeout == null ? 0 : timeoutNanos(timeout); + lock.lockInterruptibly(); + try { + while (queued.isEmpty()) { + if (closeError != null) throw closeError; + if (timeout == null) available.await(); + else if (remaining <= 0) return null; + else remaining = available.awaitNanos(remaining); + } + return queued.removeFirst(); + } finally { + lock.unlock(); + } + } + + boolean offer(RstreamStream stream) { + var accepted = false; + lock.lock(); + try { + if (closeError == null) { + queued.addLast(stream); + available.signal(); + accepted = true; + } + } finally { + lock.unlock(); + } + if (!accepted) stream.closeQuietly(); + return accepted; + } + + void close(RuntimeException error) { + ArrayList abandoned; + lock.lock(); + try { + if (closeError != null) return; + closeError = error; + abandoned = new ArrayList<>(queued); + queued.clear(); + available.signalAll(); + } finally { + lock.unlock(); + } + abandoned.forEach(RstreamStream::closeQuietly); + } + + private static long timeoutNanos(Duration timeout) { + try { + return timeout.toNanos(); + } catch (ArithmeticException ignored) { + return Long.MAX_VALUE; + } + } + } } diff --git a/src/main/java/io/rstream/ClientOptions.java b/src/main/java/io/rstream/ClientOptions.java index a7b4a58..a76cfe1 100644 --- a/src/main/java/io/rstream/ClientOptions.java +++ b/src/main/java/io/rstream/ClientOptions.java @@ -33,6 +33,13 @@ public record ClientOptions( if (operationTimeout.isZero() || operationTimeout.isNegative()) { throw new IllegalArgumentException("operationTimeout must be positive"); } + if (heartbeat + && (heartbeatInterval.compareTo(Duration.ofSeconds(1)) < 0 + || heartbeatInterval.compareTo(Duration.ofMinutes(5)) > 0 + || heartbeatInterval.toNanosPart() % 1_000_000 != 0)) { + throw new IllegalArgumentException( + "heartbeatInterval must be between 1 second and 5 minutes with millisecond precision"); + } } public static Builder builder() { diff --git a/src/main/java/io/rstream/ControlChannel.java b/src/main/java/io/rstream/ControlChannel.java index 70e54ca..8e11923 100644 --- a/src/main/java/io/rstream/ControlChannel.java +++ b/src/main/java/io/rstream/ControlChannel.java @@ -2,11 +2,19 @@ import java.io.IOException; import java.net.Socket; +import java.time.Duration; +import java.util.ArrayDeque; +import java.util.HashSet; import java.util.Map; +import java.util.Set; import java.util.UUID; +import java.util.concurrent.CancellationException; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ExecutionException; import java.util.concurrent.ExecutorService; +import java.util.concurrent.FutureTask; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.ScheduledFuture; import java.util.concurrent.TimeUnit; @@ -15,8 +23,11 @@ /** Open control channel used to create and manage tunnels. */ public final class ControlChannel implements AutoCloseable { + private static final int MAX_ACTIVE_PROXY_CONNECTIONS = 256; + private static final int MAX_QUEUED_PROXY_CONNECTIONS = 1_024; private final Socket socket; private final ResolvedClientOptions options; + private final Duration heartbeatTimeout; private final ServerDetails serverDetails; private final ExecutorService executor; private final ScheduledExecutorService scheduler; @@ -27,20 +38,28 @@ public final class ControlChannel implements AutoCloseable { private final Map tunnels = new ConcurrentHashMap<>(); private final Object writeLock = new Object(); private final Object closeLock = new Object(); + private final Object proxyLock = new Object(); + private final ArrayDeque proxyQueue = new ArrayDeque<>(); + private final Set> proxyTasks = new HashSet<>(); private final CompletableFuture done = new CompletableFuture<>(); - private ScheduledFuture heartbeat; + private volatile ScheduledFuture heartbeat; + private volatile ScheduledFuture livenessTimeout; + private volatile long heartbeatSequence; + private volatile long heartbeatAcknowledgement; private volatile boolean closing; private volatile boolean closed; ControlChannel( Socket socket, ResolvedClientOptions options, + Duration heartbeatTimeout, ServerDetails serverDetails, ExecutorService executor, ScheduledExecutorService scheduler, Function openProxyConnection) { this.socket = socket; this.options = options; + this.heartbeatTimeout = heartbeatTimeout; this.serverDetails = serverDetails; this.executor = executor; this.scheduler = scheduler; @@ -82,13 +101,15 @@ public CompletableFuture createTunnelAsync() { } public BytestreamTunnel createBytestreamTunnel(CreateTunnelOptions options) { - if (closed || closing) { - throw new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED"); - } var properties = normalizeBytestreamOptions(options); var requestId = UUID.randomUUID().toString(); var pending = new CompletableFuture(); - pendingTunnels.put(requestId, pending); + synchronized (closeLock) { + if (closed || closing) { + throw new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED"); + } + pendingTunnels.put(requestId, pending); + } var timeout = operationTimeout( pending, @@ -109,8 +130,8 @@ public void closeTunnel(String tunnelId) { if (tunnel == null || tunnel.closed()) return; CompletableFuture pending; var owner = false; - synchronized (tunnel) { - if (tunnel.closed()) return; + synchronized (closeLock) { + if (closed || closing || tunnel.closed()) return; pending = pendingCloses.get(tunnelId); if (pending == null) { pending = new CompletableFuture<>(); @@ -144,9 +165,11 @@ public CompletableFuture closeTunnelAsync(String tunnelId) { public void close() { if (closed) return; ScheduledFuture closeTimeout = null; + var sendClose = false; synchronized (closeLock) { if (!closing && !closed) { closing = true; + sendClose = true; closeTimeout = scheduler.schedule( () -> @@ -156,11 +179,13 @@ public void close() { null)), options.operationTimeout().toMillis(), TimeUnit.MILLISECONDS); - try { - write(Protocol.closeControlChannelRequest()); - } catch (RuntimeException error) { - finish(error); - } + } + } + if (sendClose) { + try { + write(Protocol.closeControlChannelRequest()); + } catch (RuntimeException error) { + finish(error); } } try { @@ -175,9 +200,21 @@ public CompletableFuture closeAsync() { } void finish(Throwable error) { - if (closed) return; - closed = true; + boolean preservePayloads; + synchronized (closeLock) { + if (closed) return; + preservePayloads = !closing && preservePayloadsAfter(error); + closed = true; + } if (heartbeat != null) heartbeat.cancel(true); + if (livenessTimeout != null) livenessTimeout.cancel(false); + Set> tasks; + synchronized (proxyLock) { + tasks = Set.copyOf(proxyTasks); + proxyTasks.clear(); + proxyQueue.clear(); + } + tasks.forEach(task -> task.cancel(true)); try { socket.close(); } catch (IOException ignored) { @@ -188,7 +225,7 @@ void finish(Throwable error) { : error; pendingTunnels.values().forEach(pending -> pending.completeExceptionally(closeError)); pendingCloses.values().forEach(pending -> pending.completeExceptionally(closeError)); - tunnels.values().forEach(tunnel -> tunnel.onClose(closeError)); + tunnels.values().forEach(tunnel -> tunnel.onClose(closeError, preservePayloads)); pendingTunnels.clear(); pendingCloses.clear(); tunnels.clear(); @@ -197,26 +234,19 @@ void finish(Throwable error) { } private void start() { - executor.submit(this::readLoop); if (options.heartbeat() && options.heartbeatInterval().toMillis() > 0) { - heartbeat = - scheduler.scheduleAtFixedRate( - () -> { - try { - if (!closed) write(Protocol.heartbeat()); - } catch (RuntimeException error) { - finish(error); - } - }, - options.heartbeatInterval().toMillis(), - options.heartbeatInterval().toMillis(), - TimeUnit.MILLISECONDS); + scheduleHeartbeat(heartbeatTimeout.isZero() ? options.heartbeatInterval().toMillis() : 0); } + if (!heartbeatTimeout.isZero()) armLivenessTimeout(); + executor.submit(this::readLoop); } private void readLoop() { try { - while (!closed) handleMessage(Protocol.readMessage(socket.getInputStream())); + while (!closed) { + handleMessage(Protocol.readMessage(socket.getInputStream())); + if (!closed && !heartbeatTimeout.isZero()) armLivenessTimeout(); + } } catch (Throwable error) { if (!closed) finish(error); } @@ -232,7 +262,11 @@ private void handleMessage(Rstream.Message message) { return; } if (message.hasProxyConnReq()) { - handleProxyConnectionRequest(message.getProxyConnReq()); + dispatchProxyConnectionRequest(message.getProxyConnReq()); + return; + } + if (message.hasHeartbeat()) { + handleHeartbeat(message.getHeartbeat()); return; } if (message.hasCloseControlChannelRsp()) finish(null); @@ -252,8 +286,17 @@ private void handleOpenTunnelResponse(Rstream.OpenTunnelRsp response) { } var properties = Protocol.tunnelPropertiesFromPb(response.getTunnelProperties()); var tunnel = new BytestreamTunnel(this, properties, executor); - tunnels.put(tunnel.id(), tunnel); - pending.complete(tunnel); + synchronized (closeLock) { + if (closed) { + tunnel.onClose( + new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED")); + pending.completeExceptionally( + new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED")); + return; + } + tunnels.put(tunnel.id(), tunnel); + pending.complete(tunnel); + } } private void handleCloseTunnelResponse(String tunnelId) { @@ -263,6 +306,81 @@ private void handleCloseTunnelResponse(String tunnelId) { if (pending != null) pending.complete(null); } + private void handleHeartbeat(Rstream.Heartbeat message) { + if (!options.heartbeat()) throw heartbeatProtocolError(); + if (heartbeatTimeout.isZero()) { + if (message.getSequence() != 0 || message.getAcknowledgement() != 0) + throw heartbeatProtocolError(); + return; + } + if (message.getSequence() != 0 + || message.getAcknowledgement() == 0 + || Long.compareUnsigned(message.getAcknowledgement(), heartbeatAcknowledgement) <= 0 + || Long.compareUnsigned(message.getAcknowledgement(), heartbeatSequence) > 0) { + throw heartbeatProtocolError(); + } + heartbeatAcknowledgement = message.getAcknowledgement(); + } + + private void dispatchProxyConnectionRequest(Rstream.ProxyConnReq request) { + FutureTask task = null; + var overloaded = false; + synchronized (proxyLock) { + if (closed) return; + if (proxyTasks.size() >= MAX_ACTIVE_PROXY_CONNECTIONS) { + if (proxyQueue.size() >= MAX_QUEUED_PROXY_CONNECTIONS) overloaded = true; + else proxyQueue.addLast(request); + } else { + task = + new FutureTask<>(() -> handleProxyConnectionRequest(request), null) { + @Override + protected void done() { + proxyTaskDone(this); + } + }; + proxyTasks.add(task); + } + } + if (overloaded) { + finish( + new RstreamException( + "Control channel proxy queue is full.", "ERR_RSTREAM_CONTROL_OVERLOAD")); + return; + } + if (task == null) return; + try { + executor.execute(task); + } catch (RejectedExecutionException error) { + synchronized (proxyLock) { + proxyTasks.remove(task); + } + if (!closed) finish(error); + } + } + + private void proxyTaskDone(FutureTask task) { + Throwable failure = null; + try { + task.get(); + } catch (CancellationException ignored) { + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + failure = error; + } catch (ExecutionException error) { + failure = error.getCause(); + } + Rstream.ProxyConnReq next = null; + synchronized (proxyLock) { + proxyTasks.remove(task); + if (!closed) next = proxyQueue.pollFirst(); + } + if (failure != null && !closed) { + finish(failure); + return; + } + if (next != null) dispatchProxyConnectionRequest(next); + } + private void handleProxyConnectionRequest(Rstream.ProxyConnReq request) { var tunnel = tunnels.get(request.getTunnelId()); if (tunnel == null) { @@ -273,6 +391,10 @@ private void handleProxyConnectionRequest(Rstream.ProxyConnReq request) { } try { var stream = openProxyConnection.apply(request); + if (closed) { + stream.closeQuietly(); + return; + } if (!tunnel.deliver(stream)) { write( Protocol.proxyConnectionResponse( @@ -288,7 +410,13 @@ private void handleProxyConnectionRequest(Rstream.ProxyConnReq request) { } private void write(Rstream.Message message) { + if (closed) { + throw new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED"); + } synchronized (writeLock) { + if (closed) { + throw new RstreamException("Control channel is closed.", "ERR_RSTREAM_CONTROL_CLOSED"); + } try { Protocol.writeMessage(socket.getOutputStream(), message); } catch (IOException error) { @@ -297,6 +425,68 @@ private void write(Rstream.Message message) { } } + private void scheduleHeartbeat(long delayMillis) { + synchronized (closeLock) { + if (closed) return; + heartbeat = + scheduler.schedule( + () -> { + try { + executor.execute(this::sendHeartbeat); + } catch (RejectedExecutionException error) { + if (!closed) finish(error); + } + }, + delayMillis, + TimeUnit.MILLISECONDS); + } + } + + private void sendHeartbeat() { + try { + if (closed) return; + if (heartbeatSequence == Long.MAX_VALUE) { + throw new ProtocolException("Heartbeat sequence exhausted.", "ERR_RSTREAM_PROTOCOL"); + } + heartbeatSequence++; + write( + heartbeatTimeout.isZero() ? Protocol.heartbeat() : Protocol.heartbeat(heartbeatSequence)); + scheduleHeartbeat(options.heartbeatInterval().toMillis()); + } catch (RuntimeException error) { + if (!closed) finish(error); + } + } + + private void armLivenessTimeout() { + synchronized (closeLock) { + if (closed) return; + if (livenessTimeout != null) livenessTimeout.cancel(false); + livenessTimeout = + scheduler.schedule( + () -> + finish( + new RstreamException( + "Control channel liveness timeout expired.", + "ERR_RSTREAM_CONTROL_LIVENESS")), + heartbeatTimeout.toMillis(), + TimeUnit.MILLISECONDS); + } + } + + private static ProtocolException heartbeatProtocolError() { + return new ProtocolException("Engine returned an invalid heartbeat.", "ERR_RSTREAM_PROTOCOL"); + } + + private static boolean preservePayloadsAfter(Throwable error) { + if (error instanceof ProtocolException) return false; + if (error instanceof RstreamException rstreamError + && rstreamError.code().equals("ERR_RSTREAM_CONTROL_LIVENESS")) return true; + for (var cause = error; cause != null; cause = cause.getCause()) { + if (cause instanceof IOException) return true; + } + return false; + } + private ScheduledFuture operationTimeout( CompletableFuture pending, Runnable cleanup, String message) { var timeout = diff --git a/src/main/java/io/rstream/Protocol.java b/src/main/java/io/rstream/Protocol.java index b7b1368..8a49a58 100644 --- a/src/main/java/io/rstream/Protocol.java +++ b/src/main/java/io/rstream/Protocol.java @@ -21,10 +21,16 @@ final class Protocol { private Protocol() {} static Rstream.Message openControlChannelRequest(String token) { - return Rstream.Message.newBuilder() - .setOpenControlChannelReq( - Rstream.OpenControlChannelReq.newBuilder().setClientDetails(clientDetails(token))) - .build(); + return openControlChannelRequest(token, null); + } + + static Rstream.Message openControlChannelRequest(String token, Integer heartbeatIntervalMs) { + var request = Rstream.OpenControlChannelReq.newBuilder().setClientDetails(clientDetails(token)); + if (heartbeatIntervalMs != null) { + request.setLiveness( + Rstream.ControlChannelLiveness.newBuilder().setHeartbeatIntervalMs(heartbeatIntervalMs)); + } + return Rstream.Message.newBuilder().setOpenControlChannelReq(request).build(); } static Rstream.Message closeControlChannelRequest() { @@ -78,6 +84,12 @@ static Rstream.Message heartbeat() { return Rstream.Message.newBuilder().setHeartbeat(Rstream.Heartbeat.newBuilder()).build(); } + static Rstream.Message heartbeat(long sequence) { + return Rstream.Message.newBuilder() + .setHeartbeat(Rstream.Heartbeat.newBuilder().setSequence(sequence)) + .build(); + } + static Rstream.ClientDetails clientDetails(String token) { var builder = Rstream.ClientDetails.newBuilder() diff --git a/src/main/java/io/rstream/RstreamClient.java b/src/main/java/io/rstream/RstreamClient.java index 7648c93..3f4aee4 100644 --- a/src/main/java/io/rstream/RstreamClient.java +++ b/src/main/java/io/rstream/RstreamClient.java @@ -18,11 +18,13 @@ /** Main rstream Java SDK client. */ public final class RstreamClient implements AutoCloseable { + private static final int MAX_HEARTBEAT_TIMEOUT_MILLIS = 900_000; private final ClientOptions options; private final RstreamTransport transport; private final ExecutorService executor; private final ScheduledExecutorService scheduler; private final Set controls = ConcurrentHashMap.newKeySet(); + private final Object lifecycleLock = new Object(); private volatile ResolvedClientOptions resolved; private volatile boolean closed; @@ -51,7 +53,11 @@ public ControlChannel connect() { try { socket = transport.dial(engine, resolvedOptions.tls(), resolvedOptions.connectTimeout()); setReadTimeout(socket, resolvedOptions.operationTimeout()); - Protocol.writeMessage(socket.getOutputStream(), Protocol.openControlChannelRequest(token)); + var heartbeatIntervalMillis = Math.toIntExact(resolvedOptions.heartbeatInterval().toMillis()); + Protocol.writeMessage( + socket.getOutputStream(), + Protocol.openControlChannelRequest( + token, resolvedOptions.heartbeat() ? heartbeatIntervalMillis : null)); var response = Protocol.readMessage(socket.getInputStream()); if (!response.hasOpenControlChannelRsp()) { throw new ProtocolException( @@ -66,15 +72,21 @@ public ControlChannel connect() { "Engine returned an empty OpenControlChannelRsp.", "ERR_RSTREAM_PROTOCOL"); } setReadTimeout(socket, Duration.ZERO); - var control = - new ControlChannel( - socket, - resolvedOptions, - Protocol.serverDetailsFromPb(payload.getOk().getServerDetails()), - executor, - scheduler, - request -> openProxyConnection(engine, resolvedOptions, request)); - controls.add(control); + ControlChannel control; + synchronized (lifecycleLock) { + ensureOpen(); + control = + new ControlChannel( + socket, + resolvedOptions, + negotiatedHeartbeatTimeout( + resolvedOptions.heartbeat(), heartbeatIntervalMillis, payload.getOk()), + Protocol.serverDetailsFromPb(payload.getOk().getServerDetails()), + executor, + scheduler, + request -> openProxyConnection(engine, resolvedOptions, request)); + controls.add(control); + } control.done().whenComplete((ignored, error) -> controls.remove(control)); return control; } catch (SocketTimeoutException error) { @@ -95,6 +107,20 @@ public CompletableFuture connectAsync() { return CompletableFuture.supplyAsync(this::connect, executor); } + private static Duration negotiatedHeartbeatTimeout( + boolean heartbeat, int heartbeatIntervalMillis, Rstream.OpenControlChannelRsp.Ok response) { + if (!response.hasLiveness()) return Duration.ZERO; + var liveness = response.getLiveness(); + if (!heartbeat + || liveness.getHeartbeatIntervalMs() != heartbeatIntervalMillis + || liveness.getHeartbeatTimeoutMs() < liveness.getHeartbeatIntervalMs() + || liveness.getHeartbeatTimeoutMs() > MAX_HEARTBEAT_TIMEOUT_MILLIS) { + throw new ProtocolException( + "Engine returned an invalid liveness policy.", "ERR_RSTREAM_PROTOCOL"); + } + return Duration.ofMillis(liveness.getHeartbeatTimeoutMs()); + } + public RstreamStream dial(String tunnel) { return dial(tunnel, DialOptions.defaults()); } @@ -111,8 +137,8 @@ public RstreamStream dial(String tunnel, DialOptions options) { var zeroRtt = options.zeroRtt() == null ? resolvedOptions.zeroRtt() : options.zeroRtt(); var socket = openStreamSocket(engine, resolvedOptions, Protocol.streamRequest(target, token, zeroRtt)); - if (!zeroRtt) { - try { + try { + if (!zeroRtt) { var response = Protocol.readMessage(socket.getInputStream()); if (!response.hasStreamRsp()) { throw new ProtocolException("Engine did not return StreamRsp.", "ERR_RSTREAM_PROTOCOL"); @@ -125,14 +151,14 @@ public RstreamStream dial(String tunnel, DialOptions options) { throw new ProtocolException( "Engine returned an empty StreamRsp.", "ERR_RSTREAM_PROTOCOL"); } - setReadTimeout(socket, Duration.ZERO); - } catch (SocketTimeoutException error) { - closeQuietly(socket); - throw operationTimeout("Timed out waiting for the private stream response.", error); - } catch (IOException | RuntimeException error) { - closeQuietly(socket); - throw runtime("Failed to dial private bytestream tunnel.", "ERR_RSTREAM_DIAL", error); } + setReadTimeout(socket, Duration.ZERO); + } catch (SocketTimeoutException error) { + closeQuietly(socket); + throw operationTimeout("Timed out waiting for the private stream response.", error); + } catch (IOException | RuntimeException error) { + closeQuietly(socket); + throw runtime("Failed to dial private bytestream tunnel.", "ERR_RSTREAM_DIAL", error); } return new RstreamStream(socket); } @@ -148,11 +174,15 @@ public CompletableFuture dialAsync(String tunnel, DialOptions opt @Override public void close() { - if (closed) return; - closed = true; + Set openControls; + synchronized (lifecycleLock) { + if (closed) return; + closed = true; + openControls = Set.copyOf(controls); + } RuntimeException failure = null; try { - for (var control : controls) { + for (var control : openControls) { try { control.close(); } catch (RuntimeException error) { @@ -199,22 +229,22 @@ private RstreamStream openProxyConnection( resolvedOptions, Protocol.proxyRequest(request.getStreamId(), token, resolvedOptions.zeroRtt()), proxyEngine.equals(engine)); - if (!resolvedOptions.zeroRtt()) { - try { + try { + if (!resolvedOptions.zeroRtt()) { var response = Protocol.readMessage(socket.getInputStream()); if (!response.hasProxyRsp()) { throw new ProtocolException("Engine did not return ProxyRsp.", "ERR_RSTREAM_PROTOCOL"); } if (response.getProxyRsp().hasError()) throw Protocol.engineErrorFromPb(response.getProxyRsp().getError()); - setReadTimeout(socket, Duration.ZERO); - } catch (SocketTimeoutException error) { - closeQuietly(socket); - throw operationTimeout("Timed out waiting for the proxy stream response.", error); - } catch (IOException | RuntimeException error) { - closeQuietly(socket); - throw runtime("Failed to open rstream proxy connection.", "ERR_RSTREAM_PROXY", error); } + setReadTimeout(socket, Duration.ZERO); + } catch (SocketTimeoutException error) { + closeQuietly(socket); + throw operationTimeout("Timed out waiting for the proxy stream response.", error); + } catch (IOException | RuntimeException error) { + closeQuietly(socket); + throw runtime("Failed to open rstream proxy connection.", "ERR_RSTREAM_PROXY", error); } return new RstreamStream(socket); } diff --git a/src/main/java/io/rstream/RstreamTransport.java b/src/main/java/io/rstream/RstreamTransport.java index d1e4735..5ac6f1b 100644 --- a/src/main/java/io/rstream/RstreamTransport.java +++ b/src/main/java/io/rstream/RstreamTransport.java @@ -17,16 +17,22 @@ SSLSocket dial(String engine, TlsOptions tls, Duration timeout, boolean useConfi var context = TlsSupport.context(tls); var peerHost = peerHost(address, tls, useConfiguredServerName); var rawSocket = new Socket(); - rawSocket.connect( - new InetSocketAddress(address.host(), address.port()), timeoutMillis(timeout)); - var socket = - (SSLSocket) - context.getSocketFactory().createSocket(rawSocket, peerHost, address.port(), true); - socket.setSoTimeout(timeoutMillis(timeout)); - TlsSupport.configure(socket, peerHost, tls); - socket.startHandshake(); - socket.setSoTimeout(0); - return socket; + SSLSocket socket = null; + try { + rawSocket.connect( + new InetSocketAddress(address.host(), address.port()), timeoutMillis(timeout)); + socket = + (SSLSocket) + context.getSocketFactory().createSocket(rawSocket, peerHost, address.port(), true); + socket.setSoTimeout(timeoutMillis(timeout)); + TlsSupport.configure(socket, peerHost, tls); + socket.startHandshake(); + socket.setSoTimeout(0); + return socket; + } catch (IOException | RuntimeException error) { + closeQuietly(socket == null ? rawSocket : socket); + throw error; + } } static String peerHost(EngineAddress address, TlsOptions tls, boolean useConfiguredServerName) { @@ -39,4 +45,11 @@ private static int timeoutMillis(Duration timeout) { if (millis > Integer.MAX_VALUE) return Integer.MAX_VALUE; return Math.max(1, (int) millis); } + + private static void closeQuietly(Socket socket) { + try { + socket.close(); + } catch (IOException ignored) { + } + } } diff --git a/src/main/proto/rstream.proto b/src/main/proto/rstream.proto index 12d23db..84aa301 100644 --- a/src/main/proto/rstream.proto +++ b/src/main/proto/rstream.proto @@ -14,7 +14,7 @@ extend google.protobuf.FieldOptions { string access = 51234; } -option (protocol_version) = "1.4.4"; +option (protocol_version) = "1.4.5"; package rstream.io_rstrm.protobuf; @@ -136,6 +136,7 @@ message TunnelProperties { message OpenControlChannelReq { ClientDetails client_details = 1; + ControlChannelLiveness liveness = 2; } // The server responds to an 'OpenControlChannelReq' message with an @@ -148,6 +149,7 @@ message OpenControlChannelRsp { message Ok { string client_id = 1; ServerDetails server_details = 2; + ControlChannelLiveness liveness = 3; } oneof payload { Ok ok = 1; @@ -263,7 +265,15 @@ message DatagramChannelClose { // Sent by client and/or server to maintain the control channel active -message Heartbeat { } +message ControlChannelLiveness { + uint32 heartbeat_interval_ms = 1; + uint32 heartbeat_timeout_ms = 2; +} + +message Heartbeat { + uint64 sequence = 1; + uint64 acknowledgement = 2; +} // Allows the server to send unsolicited messages to the client message ServerMessage { diff --git a/src/test/java/io/rstream/BytestreamTunnelTest.java b/src/test/java/io/rstream/BytestreamTunnelTest.java index ca06656..d369821 100644 --- a/src/test/java/io/rstream/BytestreamTunnelTest.java +++ b/src/test/java/io/rstream/BytestreamTunnelTest.java @@ -3,8 +3,15 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import java.net.InetAddress; +import java.net.ServerSocket; +import java.net.Socket; import java.time.Duration; +import java.util.concurrent.ExecutionException; import java.util.concurrent.Executors; +import java.util.concurrent.LinkedBlockingQueue; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; import org.junit.jupiter.api.Test; final class BytestreamTunnelTest { @@ -104,4 +111,98 @@ void constructorRejectsMissingTunnelId() { executor.shutdownNow(); } } + + @Test + void hardCloseClosesAForwarderBeforeItsTaskStarts() throws Exception { + var executor = singleWorkerExecutor(); + try (var peerListener = loopbackListener(); + var targetListener = loopbackListener(); + var streamSocket = new Socket(peerListener.getInetAddress(), peerListener.getLocalPort()); + var peer = peerListener.accept()) { + var tunnel = tunnel(executor); + var forwarding = tunnel.forwardTo("127.0.0.1", targetListener.getLocalPort()); + assertThat(tunnel.deliver(new RstreamStream(streamSocket))).isTrue(); + awaitQueuedTasks(executor, 1); + tunnel.onClose(null, false); + peer.setSoTimeout(1_000); + assertThat(peer.getInputStream().read()).isEqualTo(-1); + forwarding.get(2, TimeUnit.SECONDS); + assertThat(executor.getQueue()).isEmpty(); + } finally { + executor.shutdownNow(); + } + } + + @Test + void softClosePreservesAForwarderBeforeItsTaskStarts() throws Exception { + var executor = singleWorkerExecutor(); + try (var peerListener = loopbackListener(); + var targetListener = loopbackListener(); + var streamSocket = new Socket(peerListener.getInetAddress(), peerListener.getLocalPort()); + var peer = peerListener.accept()) { + var tunnel = tunnel(executor); + var forwarding = tunnel.forwardTo("127.0.0.1", targetListener.getLocalPort()); + assertThat(tunnel.deliver(new RstreamStream(streamSocket))).isTrue(); + awaitQueuedTasks(executor, 1); + tunnel.onClose( + new RstreamException("Control transport lost.", "ERR_RSTREAM_CONTROL_LIVENESS"), true); + try (var target = targetListener.accept()) { + var bytes = "survives".getBytes(java.nio.charset.StandardCharsets.UTF_8); + peer.getOutputStream().write(bytes); + peer.getOutputStream().flush(); + assertThat(target.getInputStream().readNBytes(bytes.length)).isEqualTo(bytes); + target.getOutputStream().write(bytes); + target.getOutputStream().flush(); + assertThat(peer.getInputStream().readNBytes(bytes.length)).isEqualTo(bytes); + } + forwarding.get(2, TimeUnit.SECONDS); + } finally { + executor.shutdownNow(); + } + } + + @Test + void hardCloseCannotLeakALiveStreamThroughTheAcceptanceRace() throws Exception { + var executor = Executors.newCachedThreadPool(); + try { + for (var iteration = 0; iteration < 200; iteration++) { + var socket = new Socket(); + var tunnel = + new BytestreamTunnel( + null, TunnelProperties.builder().id("tun_" + iteration).build(), executor); + var acceptance = tunnel.acceptAsync(); + assertThat(tunnel.deliver(new RstreamStream(socket))).isTrue(); + tunnel.onClose(null, false); + try { + assertThat(acceptance.get(2, TimeUnit.SECONDS).socket().isClosed()).isTrue(); + } catch (ExecutionException error) { + assertThat(error.getCause()).isInstanceOf(RstreamException.class); + } + } + } finally { + executor.shutdownNow(); + } + } + + private static BytestreamTunnel tunnel(ThreadPoolExecutor executor) { + return new BytestreamTunnel( + null, TunnelProperties.builder().id("tun_1").name("private-api").build(), executor); + } + + private static ServerSocket loopbackListener() throws java.io.IOException { + return new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + } + + private static ThreadPoolExecutor singleWorkerExecutor() { + return new ThreadPoolExecutor(1, 1, 0, TimeUnit.MILLISECONDS, new LinkedBlockingQueue<>()); + } + + private static void awaitQueuedTasks(ThreadPoolExecutor executor, int expected) + throws InterruptedException { + var deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(2); + while (executor.getQueue().size() != expected && System.nanoTime() < deadline) { + Thread.sleep(1); + } + assertThat(executor.getQueue()).hasSize(expected); + } } diff --git a/src/test/java/io/rstream/ConfigResolverTest.java b/src/test/java/io/rstream/ConfigResolverTest.java index 9eb1166..6a02b28 100644 --- a/src/test/java/io/rstream/ConfigResolverTest.java +++ b/src/test/java/io/rstream/ConfigResolverTest.java @@ -15,6 +15,23 @@ import org.junit.jupiter.api.io.TempDir; final class ConfigResolverTest { + @Test + void heartbeatIntervalMustUseSupportedMillisecondBounds() { + assertThatThrownBy( + () -> ClientOptions.builder().heartbeatInterval(Duration.ofMillis(999)).build()) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("heartbeatInterval"); + assertThatThrownBy( + () -> ClientOptions.builder().heartbeatInterval(Duration.ofMillis(300_001)).build()) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("heartbeatInterval"); + assertThatThrownBy( + () -> + ClientOptions.builder().heartbeatInterval(Duration.ofNanos(1_000_500_000)).build()) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("heartbeatInterval"); + } + @TempDir Path temp; @Test diff --git a/src/test/java/io/rstream/ProtocolTest.java b/src/test/java/io/rstream/ProtocolTest.java index 3a87ae7..8dbe817 100644 --- a/src/test/java/io/rstream/ProtocolTest.java +++ b/src/test/java/io/rstream/ProtocolTest.java @@ -70,6 +70,16 @@ void streamRequestContainsTargetAuthAndZeroRtt() { assertThat(request.getZeroRtt().getValue()).isFalse(); } + @Test + void controlLivenessMessagesPreserveIntervalAndSequence() { + var request = Protocol.openControlChannelRequest(null, 1_250).getOpenControlChannelReq(); + var heartbeat = Protocol.heartbeat(42).getHeartbeat(); + assertThat(request.getLiveness().getHeartbeatIntervalMs()).isEqualTo(1_250); + assertThat(request.getLiveness().getHeartbeatTimeoutMs()).isZero(); + assertThat(heartbeat.getSequence()).isEqualTo(42); + assertThat(heartbeat.getAcknowledgement()).isZero(); + } + @Test void proxyRequestContainsStreamAuthAndZeroRtt() { var request = Protocol.proxyRequest("stream_123", null, true).getProxyReq(); diff --git a/src/test/java/io/rstream/RstreamTransportTest.java b/src/test/java/io/rstream/RstreamTransportTest.java index 4c74cc4..58e7768 100644 --- a/src/test/java/io/rstream/RstreamTransportTest.java +++ b/src/test/java/io/rstream/RstreamTransportTest.java @@ -3,12 +3,15 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import java.io.OutputStream; import java.net.InetAddress; import java.net.ServerSocket; +import java.net.Socket; import java.time.Duration; -import java.util.concurrent.CountDownLatch; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executors; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.Test; final class RstreamTransportTest { @@ -23,15 +26,16 @@ void redirectedEngineUsesItsHostnameForTls() { @Test void tlsHandshakeUsesConfiguredTimeout() throws Exception { - var accepted = new CountDownLatch(1); + var peerClosed = new CompletableFuture(); + var acceptedSocket = new AtomicReference(); var executor = Executors.newSingleThreadExecutor(); try (var listener = new ServerSocket(0, 1, InetAddress.getLoopbackAddress())) { executor.submit( () -> { try (var socket = listener.accept()) { - socket.setSoTimeout(2_000); - accepted.countDown(); - Thread.sleep(2_000); + acceptedSocket.set(socket); + socket.getInputStream().transferTo(OutputStream.nullOutputStream()); + peerClosed.complete(null); } catch (Exception ignored) { } }); @@ -44,8 +48,10 @@ void tlsHandshakeUsesConfiguredTimeout() throws Exception { TlsOptions.builder().insecureSkipVerify(true).build(), Duration.ofMillis(100))) .isInstanceOf(java.net.SocketTimeoutException.class); - accepted.await(1, TimeUnit.SECONDS); + peerClosed.get(1, TimeUnit.SECONDS); } finally { + var socket = acceptedSocket.get(); + if (socket != null) socket.close(); executor.shutdownNow(); } } diff --git a/src/test/java/io/rstream/RuntimeFakeEngineIT.java b/src/test/java/io/rstream/RuntimeFakeEngineIT.java index 54bb9cd..865092e 100644 --- a/src/test/java/io/rstream/RuntimeFakeEngineIT.java +++ b/src/test/java/io/rstream/RuntimeFakeEngineIT.java @@ -7,6 +7,8 @@ import java.io.Closeable; import java.io.IOException; import java.net.InetAddress; +import java.net.ServerSocket; +import java.net.Socket; import java.nio.charset.StandardCharsets; import java.nio.file.Path; import java.security.KeyStore; @@ -16,6 +18,7 @@ import java.util.Map; import java.util.concurrent.BlockingQueue; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutionException; import java.util.concurrent.LinkedBlockingQueue; import java.util.concurrent.TimeUnit; @@ -65,6 +68,292 @@ void createTunnelSendsNormalizedPropertiesAndCloses() throws Exception { } } + @Test + void negotiatedLivenessToleratesDelayedAcknowledgement() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 1_000, true, 800, 1, 0); + try (var control = client.connect()) { + var heartbeat = engine.heartbeats.poll(2, TimeUnit.SECONDS); + assertThat(heartbeat).isNotNull(); + assertThat(engine.openControlRequests.peek().getLiveness().getHeartbeatIntervalMs()) + .isEqualTo(1_000); + assertThat(heartbeat.getSequence()).isEqualTo(1); + Thread.sleep(1_100); + assertThat(control.closed()).isFalse(); + } + } + } + + @Test + void negotiatedLivenessExpiresWhenAcknowledgementsStop() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 1_000, false, 0, 1, 0); + var control = client.connect(); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)) + .isEqualTo("ERR_RSTREAM_CONTROL_LIVENESS"); + } + } + + @Test + void acceptedStreamSurvivesLivenessTimeout() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 1_000, false, 0, 1, 0); + var control = client.connect(); + var tunnel = control.createTunnel(); + engine.sendProxyConnection(tunnel.id(), "draining-accepted", "stream-secret"); + try (var stream = tunnel.accept(Duration.ofSeconds(2))) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(stream.socket(), "before")).isEqualTo("before"); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)) + .isEqualTo("ERR_RSTREAM_CONTROL_LIVENESS"); + assertThat(tunnel.closed()).isTrue(); + assertThat(roundTrip(stream.socket(), "after")).isEqualTo("after"); + } + } + } + + @Test + void forwardedStreamSurvivesLivenessTimeout() throws Exception { + try (var local = new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 1_000, false, 0, 1, 0); + var control = client.connect(); + var tunnel = control.createTunnel(); + var forwarding = tunnel.forwardTo("127.0.0.1", local.getLocalPort()); + var accepted = CompletableFuture.supplyAsync(() -> accept(local)); + engine.sendProxyConnection(tunnel.id(), "draining-forward", "stream-secret"); + try (var localStream = accepted.get(2, TimeUnit.SECONDS)) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(localStream, "before")).isEqualTo("before"); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)) + .isEqualTo("ERR_RSTREAM_CONTROL_LIVENESS"); + assertThat(roundTrip(localStream, "after")).isEqualTo("after"); + } + forwarding.get(2, TimeUnit.SECONDS); + } + } + + @Test + void explicitControlCloseStopsForwardedStreamLocally() throws Exception { + try (var local = new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + var forwarding = tunnel.forwardTo("127.0.0.1", local.getLocalPort()); + var accepted = CompletableFuture.supplyAsync(() -> accept(local)); + engine.sendProxyConnection(tunnel.id(), "hard-close-forward", "stream-secret"); + try (var localStream = accepted.get(2, TimeUnit.SECONDS)) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(localStream, "before")).isEqualTo("before"); + control.close(); + localStream.setSoTimeout(500); + assertThat(localStream.getInputStream().read()).isEqualTo(-1); + } + forwarding.get(2, TimeUnit.SECONDS); + } + } + + @Test + void explicitControlCloseStopsAcceptedStreamLocally() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + engine.sendProxyConnection(tunnel.id(), "hard-close-accepted", "stream-secret"); + try (var stream = tunnel.accept(Duration.ofSeconds(2))) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(stream.socket(), "before")).isEqualTo("before"); + control.close(); + assertThat(stream.socket().isClosed()).isTrue(); + } + } + } + + @Test + void localHardCloseAfterLivenessTimeoutStopsAcceptedStream() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 1_000, false, 0, 1, 0); + var control = client.connect(); + var tunnel = control.createTunnel(); + engine.sendProxyConnection(tunnel.id(), "soft-then-hard-accepted", "stream-secret"); + try (var stream = tunnel.accept(Duration.ofSeconds(2))) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(stream.socket(), "before")).isEqualTo("before"); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)) + .isEqualTo("ERR_RSTREAM_CONTROL_LIVENESS"); + assertThat(roundTrip(stream.socket(), "after-soft-close")).isEqualTo("after-soft-close"); + tunnel.close(); + assertThat(stream.socket().isClosed()).isTrue(); + } + } + } + + @Test + void malformedControlFrameStopsForwardedStreamLocally() throws Exception { + try (var local = new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + var forwarding = tunnel.forwardTo("127.0.0.1", local.getLocalPort()); + var accepted = CompletableFuture.supplyAsync(() -> accept(local)); + engine.sendProxyConnection(tunnel.id(), "protocol-failure-forward", "stream-secret"); + try (var localStream = accepted.get(2, TimeUnit.SECONDS)) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(localStream, "before")).isEqualTo("before"); + engine.sendMalformedControlFrame(); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)).isEqualTo("ERR_RSTREAM_PROTOCOL"); + localStream.setSoTimeout(500); + assertThat(localStream.getInputStream().read()).isEqualTo(-1); + } + forwarding.get(2, TimeUnit.SECONDS); + } + } + + @Test + void zeroRttDirectStreamDoesNotInheritOperationTimeout() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = timeoutClient(engine, true)) { + CompletableFuture blockedRead; + try (var stream = client.dial("private-api")) { + blockedRead = + CompletableFuture.supplyAsync( + () -> { + try { + return stream.inputStream().read(); + } catch (IOException error) { + throw new RstreamException("Test read failed.", "ERR_TEST_READ", error); + } + }); + Thread.sleep(250); + assertThat(blockedRead).isNotDone(); + } + blockedRead.handle((value, error) -> null).get(2, TimeUnit.SECONDS); + } + } + + @Test + void zeroRttForwarderDoesNotInheritOperationTimeout() throws Exception { + try (var local = new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + var engine = FakeEngine.start(temp); + var client = timeoutClient(engine, true); + var control = client.connect()) { + var tunnel = control.createTunnel(); + var forwarding = tunnel.forwardTo("127.0.0.1", local.getLocalPort()); + var accepted = CompletableFuture.supplyAsync(() -> accept(local)); + engine.sendProxyConnection(tunnel.id(), "zero-rtt-timeout-forward", "stream-secret"); + try (var localStream = accepted.get(2, TimeUnit.SECONDS)) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + assertThat(roundTrip(localStream, "before")).isEqualTo("before"); + Thread.sleep(250); + assertThat(roundTrip(localStream, "after")).isEqualTo("after"); + } + tunnel.close(); + forwarding.get(2, TimeUnit.SECONDS); + } + } + + @Test + void acceptedStreamSurvivesUnexpectedControlTransportEof() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + engine.sendProxyConnection(tunnel.id(), "eof-accepted", "stream-secret"); + try (var stream = tunnel.accept(Duration.ofSeconds(2))) { + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + engine.closeControlSocket(); + assertThatThrownBy(() -> control.done().get(2, TimeUnit.SECONDS)) + .isInstanceOf(ExecutionException.class); + assertThat(roundTrip(stream.socket(), "survives")).isEqualTo("survives"); + } + } + } + + @Test + void negotiatedLivenessRejectsFutureAcknowledgement() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 60_000, true, 0, 1, 1); + var control = client.connect(); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)).isEqualTo("ERR_RSTREAM_PROTOCOL"); + } + } + + @Test + void negotiatedLivenessRejectsReplayedAcknowledgement() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 60_000, true, 0, 1, 0); + engine.duplicateHeartbeatAcknowledgement = true; + var control = client.connect(); + assertThat(failureCode(control, 2, TimeUnit.SECONDS)).isEqualTo("ERR_RSTREAM_PROTOCOL"); + } + } + + @Test + void negotiatedLivenessToleratesDroppedHeartbeat() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(1_000, 2_500, true, 0, 2, 0); + try (var control = client.connect()) { + for (var sequence = 1L; sequence <= 4L; sequence++) { + var heartbeat = engine.heartbeats.poll(3, TimeUnit.SECONDS); + assertThat(heartbeat).isNotNull(); + assertThat(heartbeat.getSequence()).isEqualTo(sequence); + } + assertThat(control.closed()).isFalse(); + } + } + } + + @Test + void connectRejectsInvalidServerLivenessPolicies() throws Exception { + for (var policy : + List.of(new int[] {2_000, 60_000}, new int[] {1_000, 999}, new int[] {1_000, 900_001})) { + try (var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofSeconds(1))) { + engine.configureLiveness(policy[0], policy[1], false, 0, 1, 0); + assertThatThrownBy(client::connect).isInstanceOf(ProtocolException.class); + } + } + try (var engine = FakeEngine.start(temp); + var client = client(engine)) { + engine.configureLiveness(1_000, 60_000, false, 0, 1, 0); + assertThatThrownBy(client::connect).isInstanceOf(ProtocolException.class); + } + } + + @Test + void livenessIsNotStarvedByStalledProxyTlsHandshake() throws Exception { + try (var blackhole = new ServerSocket(0, 50, InetAddress.getLoopbackAddress()); + var engine = FakeEngine.start(temp); + var client = heartbeatClient(engine, Duration.ofMillis(500))) { + engine.configureLiveness(1_000, 1_500, true, 0, 1, 0); + try (var control = client.connect()) { + var tunnel = control.createTunnel(); + var accepted = CompletableFuture.supplyAsync(() -> accept(blackhole)); + engine.sendProxyConnection( + tunnel.id(), + "blocked-stream", + "stream-secret", + "127.0.0.1:" + blackhole.getLocalPort()); + try (var blackholeSocket = accepted.get(1, TimeUnit.SECONDS)) { + assertThat(blackholeSocket.isClosed()).isFalse(); + Thread.sleep(1_900); + assertThat(control.closed()).isFalse(); + assertThat(control.createTunnel().id()).isNotBlank(); + } + } + } + } + @ParameterizedTest @ValueSource(booleans = {false, true}) void dialPrivateBytestreamByNameAndId(boolean zeroRtt) throws Exception { @@ -110,6 +399,39 @@ void proxyConnectionIsDeliveredToTunnel() throws Exception { } } + @Test + void controlCloseReleasesAllConcurrentAcceptWaiters() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + var accepts = IntStream.range(0, 8).mapToObj(ignored -> tunnel.acceptAsync()).toList(); + Thread.sleep(50); + + control.close(); + + assertThat(accepts).allMatch(CompletableFuture::isDone); + assertThat(accepts).allMatch(CompletableFuture::isCompletedExceptionally); + } + } + + @Test + void controlCloseClosesUnacceptedProxyStream() throws Exception { + try (var engine = FakeEngine.start(temp); + var client = client(engine)) { + var control = client.connect(); + var tunnel = control.createTunnel(); + engine.sendProxyConnection(tunnel.id(), "unaccepted-stream", "stream-secret"); + assertThat(engine.proxyConnectionResponses.poll(2, TimeUnit.SECONDS)).isNotNull(); + var proxyClosed = engine.proxyClosures.poll(2, TimeUnit.SECONDS); + assertThat(proxyClosed).isNotNull(); + + control.close(); + + proxyClosed.get(1, TimeUnit.SECONDS); + } + } + @Test void proxyConnectionCanDialIngressEngine() throws Exception { try (var owner = FakeEngine.start(temp); @@ -322,6 +644,24 @@ void clientCloseClosesOpenControlChannels() throws Exception { } } + @Test + void clientCloseWinsRaceWithControlRegistration() throws Exception { + try (var engine = FakeEngine.start(temp)) { + var client = client(engine); + engine.pauseControlResponse(); + var connecting = CompletableFuture.supplyAsync(client::connect); + assertThat(engine.openControlRequests.poll(2, TimeUnit.SECONDS)).isNotNull(); + + client.close(); + engine.resumeControlResponse(); + + assertThatThrownBy(() -> connecting.get(2, TimeUnit.SECONDS)) + .isInstanceOf(ExecutionException.class) + .hasCauseInstanceOf(RstreamException.class) + .hasRootCauseMessage("rstream client is closed."); + } + } + @Test void controlOpenTimeoutIsBounded() throws Exception { try (var engine = FakeEngine.start(temp); @@ -638,6 +978,39 @@ private static RstreamClient client(FakeEngine engine) { return client(engine, true); } + private static RstreamClient heartbeatClient(FakeEngine engine, Duration operationTimeout) { + return RstreamClient.fromEnv( + ClientOptions.builder() + .engine(engine.address()) + .readConfigFile(false) + .noToken(true) + .heartbeat(true) + .heartbeatInterval(Duration.ofSeconds(1)) + .connectTimeout(Duration.ofSeconds(5)) + .operationTimeout(operationTimeout) + .tls(TlsOptions.builder().insecureSkipVerify(true).build()) + .build()); + } + + private static String failureCode(ControlChannel control, long timeout, TimeUnit timeUnit) + throws Exception { + try { + control.done().get(timeout, timeUnit); + throw new AssertionError("control channel completed without an error"); + } catch (ExecutionException error) { + assertThat(error.getCause()).isInstanceOf(RstreamException.class); + return ((RstreamException) error.getCause()).code(); + } + } + + private static Socket accept(ServerSocket server) { + try { + return server.accept(); + } catch (IOException error) { + throw new RstreamException("Test blackhole accept failed.", "ERR_TEST_ENGINE", error); + } + } + private static RstreamClient timeoutClient(FakeEngine engine) { return timeoutClient(engine, true); } @@ -684,6 +1057,13 @@ private static String dialEcho( } } + private static String roundTrip(Socket socket, String payload) throws IOException { + var bytes = payload.getBytes(StandardCharsets.UTF_8); + socket.getOutputStream().write(bytes); + socket.getOutputStream().flush(); + return new String(socket.getInputStream().readNBytes(bytes.length), StandardCharsets.UTF_8); + } + private static List drain(BlockingQueue queue, int count) throws InterruptedException { var values = new ArrayList(); for (var index = 0; index < count; index++) { @@ -712,13 +1092,26 @@ private static final class FakeEngine implements Closeable { private volatile boolean nextStreamHang; private volatile boolean nextStreamEmptyResponse; private volatile boolean nextProxyHang; + private volatile Rstream.ControlChannelLiveness liveness; + private volatile boolean acknowledgeHeartbeats; + private volatile long heartbeatAcknowledgementDelayMillis; + private volatile int heartbeatAcknowledgementEvery = 1; + private volatile long heartbeatAcknowledgementOffset; + private volatile boolean duplicateHeartbeatAcknowledgement; + private volatile CountDownLatch controlResponseGate; private int tunnelCounter; + private int heartbeatCount; + private final BlockingQueue openControlRequests = + new LinkedBlockingQueue<>(); + private final BlockingQueue heartbeats = new LinkedBlockingQueue<>(); private final BlockingQueue openTunnelRequests = new LinkedBlockingQueue<>(); private final BlockingQueue closeTunnelRequests = new LinkedBlockingQueue<>(); private final BlockingQueue streamRequests = new LinkedBlockingQueue<>(); private final BlockingQueue proxyRequests = new LinkedBlockingQueue<>(); + private final BlockingQueue> proxyClosures = + new LinkedBlockingQueue<>(); private final BlockingQueue proxyConnectionResponses = new LinkedBlockingQueue<>(); @@ -746,6 +1139,32 @@ Path certificatePath() { return certificatePath; } + void configureLiveness( + int intervalMillis, + int timeoutMillis, + boolean acknowledge, + long acknowledgementDelayMillis, + int acknowledgementEvery, + long acknowledgementOffset) { + liveness = + Rstream.ControlChannelLiveness.newBuilder() + .setHeartbeatIntervalMs(intervalMillis) + .setHeartbeatTimeoutMs(timeoutMillis) + .build(); + acknowledgeHeartbeats = acknowledge; + heartbeatAcknowledgementDelayMillis = acknowledgementDelayMillis; + heartbeatAcknowledgementEvery = acknowledgementEvery; + heartbeatAcknowledgementOffset = acknowledgementOffset; + } + + void pauseControlResponse() { + controlResponseGate = new CountDownLatch(1); + } + + void resumeControlResponse() { + controlResponseGate.countDown(); + } + void sendProxyConnection(String tunnelId, String streamId, String secret) { sendProxyConnection(tunnelId, streamId, secret, null); } @@ -763,6 +1182,17 @@ void closeControlSocket() throws IOException { controlSocket.close(); } + void sendMalformedControlFrame() { + synchronized (controlWriteLock) { + try { + controlSocket.getOutputStream().write(new byte[] {0, 0, 0, 1, (byte) 0xff}); + controlSocket.getOutputStream().flush(); + } catch (IOException error) { + throw new RstreamException("Fake engine control write failed.", "ERR_TEST_ENGINE", error); + } + } + } + @Override public void close() throws IOException { listener.close(); @@ -786,6 +1216,7 @@ private void handle(SSLSocket socket) { try { var message = Protocol.readMessage(socket.getInputStream()); if (message.hasOpenControlChannelReq()) { + openControlRequests.offer(message.getOpenControlChannelReq()); handleControl(socket); return; } @@ -814,22 +1245,32 @@ private void handleControl(SSLSocket socket) throws IOException { nextControlHang = false; return; } + var responseGate = controlResponseGate; + if (responseGate != null) { + try { + responseGate.await(); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + return; + } + } + var ok = + Rstream.OpenControlChannelRsp.Ok.newBuilder() + .setClientId("cli_1") + .setServerDetails( + Rstream.ServerDetails.newBuilder() + .setAgent(StringValue.newBuilder().setValue("fake-engine"))); + if (liveness != null) ok.setLiveness(liveness); writeControl( Rstream.Message.newBuilder() - .setOpenControlChannelRsp( - Rstream.OpenControlChannelRsp.newBuilder() - .setOk( - Rstream.OpenControlChannelRsp.Ok.newBuilder() - .setClientId("cli_1") - .setServerDetails( - Rstream.ServerDetails.newBuilder() - .setAgent(StringValue.newBuilder().setValue("fake-engine"))))) + .setOpenControlChannelRsp(Rstream.OpenControlChannelRsp.newBuilder().setOk(ok)) .build()); while (!socket.isClosed()) { var message = Protocol.readMessage(socket.getInputStream()); if (message.hasOpenTunnelReq()) handleOpenTunnel(message.getOpenTunnelReq()); if (message.hasCloseTunnelReq()) handleCloseTunnel(message.getCloseTunnelReq()); if (message.hasProxyConnRsp()) proxyConnectionResponses.offer(message.getProxyConnRsp()); + if (message.hasHeartbeat()) handleHeartbeat(message.getHeartbeat()); if (message.hasCloseControlChannelReq()) { if (nextCloseControlHang) { nextCloseControlHang = false; @@ -844,6 +1285,35 @@ private void handleControl(SSLSocket socket) throws IOException { } } + private void handleHeartbeat(Rstream.Heartbeat heartbeat) { + heartbeats.offer(heartbeat); + heartbeatCount++; + if (!acknowledgeHeartbeats || heartbeatCount % heartbeatAcknowledgementEvery != 0) return; + if (heartbeatCount == 1 && heartbeatAcknowledgementDelayMillis > 0) { + try { + Thread.sleep(heartbeatAcknowledgementDelayMillis); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + return; + } + } + writeControl( + Rstream.Message.newBuilder() + .setHeartbeat( + Rstream.Heartbeat.newBuilder() + .setAcknowledgement(heartbeat.getSequence() + heartbeatAcknowledgementOffset)) + .build()); + if (duplicateHeartbeatAcknowledgement) { + writeControl( + Rstream.Message.newBuilder() + .setHeartbeat( + Rstream.Heartbeat.newBuilder() + .setAcknowledgement( + heartbeat.getSequence() + heartbeatAcknowledgementOffset)) + .build()); + } + } + private void handleOpenTunnel(Rstream.OpenTunnelReq request) { openTunnelRequests.offer(request); var builder = Rstream.OpenTunnelRsp.newBuilder().setRequestId(request.getRequestId()); @@ -929,7 +1399,13 @@ private void handleProxy(SSLSocket socket, Rstream.ProxyReq request) throws IOEx socket.getOutputStream(), Rstream.Message.newBuilder().setProxyRsp(Rstream.ProxyRsp.newBuilder()).build()); } - echo(socket); + var closed = new CompletableFuture(); + proxyClosures.offer(closed); + try { + echo(socket); + } finally { + closed.complete(null); + } } private void writeControl(Rstream.Message message) { diff --git a/src/test/java/io/rstream/RuntimeRealEngineIT.java b/src/test/java/io/rstream/RuntimeRealEngineIT.java index 2d3e6a7..e468301 100644 --- a/src/test/java/io/rstream/RuntimeRealEngineIT.java +++ b/src/test/java/io/rstream/RuntimeRealEngineIT.java @@ -120,6 +120,14 @@ void connectAsyncOpensAControlChannel() throws Exception { var control = client.connectAsync().get(20, TimeUnit.SECONDS)) { assertThat(control.serverDetails()).isNotNull(); assertThat(control.closed()).isFalse(); + Thread.sleep(2200); + var tunnel = + control.createTunnel( + CreateTunnelOptions.builder() + .name("java-heartbeat-" + UUID.randomUUID().toString().substring(0, 8)) + .publish(false) + .build()); + control.closeTunnel(tunnel.id()); } } @@ -339,7 +347,7 @@ private static String envOrDefault(String key, String fallback) { } private static RstreamClient client() { - var builder = ClientOptions.builder().heartbeat(false); + var builder = ClientOptions.builder().heartbeatInterval(Duration.ofSeconds(1)); var engine = System.getenv("RSTREAM_JAVA_E2E_ENGINE"); if (engine != null && !engine.isBlank()) builder.engine(engine); if ("1".equals(System.getenv("RSTREAM_JAVA_E2E_NO_TOKEN"))) builder.noToken(true);