diff --git a/crates/openshell-cli/src/run.rs b/crates/openshell-cli/src/run.rs index 02c806c421..2a23e91161 100644 --- a/crates/openshell-cli/src/run.rs +++ b/crates/openshell-cli/src/run.rs @@ -2232,7 +2232,6 @@ async fn forward_one_tcp_connection( service_id: String, authorization_token: String, ) -> std::result::Result<(), ForwardTcpConnectionError> { - use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio_stream::wrappers::ReceiverStream; let (tx, rx) = tokio::sync::mpsc::channel::(16); @@ -2247,13 +2246,14 @@ async fn forward_one_tcp_connection( port: u32::from(target_port), })), authorization_token, + capabilities: openshell_core::stream_lifecycle::capabilities(), }, )), }) .await .map_err(|_| ForwardTcpConnectionError::transient("failed to initialize forward stream"))?; - let mut response = match client.forward_tcp(ReceiverStream::new(rx)).await { + let response = match client.forward_tcp(ReceiverStream::new(rx)).await { Ok(response) => response.into_inner(), Err(status) => { let err = ForwardTcpConnectionError::from_status(status); @@ -2262,51 +2262,10 @@ async fn forward_one_tcp_connection( } }; - let (mut local_read, mut local_write) = socket.into_split(); - - let to_gateway = tokio::spawn(async move { - let mut buf = vec![0u8; 64 * 1024]; - loop { - let n = local_read.read(&mut buf).await?; - if n == 0 { - break; - } - if tx - .send(TcpForwardFrame { - payload: Some(openshell_core::proto::tcp_forward_frame::Payload::Data( - buf[..n].to_vec(), - )), - }) - .await - .is_err() - { - break; - } - } - Ok::<(), std::io::Error>(()) - }); - - while let Some(frame) = response - .message() + let (local_read, local_write) = socket.into_split(); + openshell_core::stream_lifecycle::client(response, local_read, local_write, tx, true) .await - .map_err(ForwardTcpConnectionError::from_status)? - { - let Some(openshell_core::proto::tcp_forward_frame::Payload::Data(data)) = frame.payload - else { - continue; - }; - if data.is_empty() { - continue; - } - local_write - .write_all(&data) - .await - .map_err(|err| ForwardTcpConnectionError::transient(err.to_string()))?; - } - - let _ = local_write.shutdown().await; - to_gateway.abort(); - Ok(()) + .map_err(ForwardTcpConnectionError::from_status) } async fn drain_and_shutdown_local_socket(mut socket: tokio::net::TcpStream) { diff --git a/crates/openshell-cli/src/ssh.rs b/crates/openshell-cli/src/ssh.rs index 8d818062ca..f60a1cb28f 100644 --- a/crates/openshell-cli/src/ssh.rs +++ b/crates/openshell-cli/src/ssh.rs @@ -25,7 +25,6 @@ use std::os::unix::process::CommandExt; use std::path::{Path, PathBuf}; use std::process::{Command, ExitStatus, Stdio}; use std::time::{Duration, Instant}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpStream; use tokio::process::{Child, Command as TokioCommand}; use tokio_stream::wrappers::ReceiverStream; @@ -1914,64 +1913,28 @@ pub async fn sandbox_ssh_proxy( service_id: format!("ssh-proxy:{sandbox_name}"), target: Some(tcp_forward_init::Target::Ssh(SshRelayTarget {})), authorization_token: token.to_string(), + capabilities: Vec::new(), }, )), }) .await .map_err(|_| miette::miette!("failed to initialize SSH forward stream"))?; - let mut response = client + let response = client .forward_tcp(ReceiverStream::new(rx)) .await .into_diagnostic()? .into_inner(); - let stdin = tokio::io::stdin(); - let stdout = tokio::io::stdout(); - - let to_remote = tokio::spawn(async move { - let mut stdin = stdin; - let mut buf = vec![0u8; 64 * 1024]; - while let Ok(n) = stdin.read(&mut buf).await { - if n == 0 { - break; - } - if tx - .send(TcpForwardFrame { - payload: Some(openshell_core::proto::tcp_forward_frame::Payload::Data( - buf[..n].to_vec(), - )), - }) - .await - .is_err() - { - break; - } - } - }); - let from_remote = tokio::spawn(async move { - let mut stdout = stdout; - loop { - let Ok(Some(frame)) = response.message().await else { - break; - }; - let Some(openshell_core::proto::tcp_forward_frame::Payload::Data(data)) = frame.payload - else { - continue; - }; - if data.is_empty() { - continue; - } - if stdout.write_all(&data).await.is_err() { - break; - } - let _ = stdout.flush().await; - } - }); - let _ = from_remote.await; - to_remote.abort(); - - Ok(()) + openshell_core::stream_lifecycle::client( + response, + tokio::io::stdin(), + tokio::io::stdout(), + tx, + false, + ) + .await + .into_diagnostic() } fn grpc_server_from_ssh_gateway_url(gateway_url: &str) -> Result { diff --git a/crates/openshell-core/src/lib.rs b/crates/openshell-core/src/lib.rs index 472175f511..b41c026fbe 100644 --- a/crates/openshell-core/src/lib.rs +++ b/crates/openshell-core/src/lib.rs @@ -54,6 +54,7 @@ pub mod secrets; pub mod settings; pub mod shell; pub mod spiffe; +pub mod stream_lifecycle; pub mod telemetry; pub mod time; pub mod transport_errors; diff --git a/crates/openshell-core/src/stream_lifecycle.rs b/crates/openshell-core/src/stream_lifecycle.rs new file mode 100644 index 0000000000..ea56acfe89 --- /dev/null +++ b/crates/openshell-core/src/stream_lifecycle.rs @@ -0,0 +1,1001 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Directional byte-stream closure and abort-preserving relay transport. + +use crate::proto::{RelayClose, RelayCloseCode}; +use std::io; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll, Waker}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream, ReadBuf}; +use tokio::sync::{mpsc, watch}; +use tokio_stream::{Stream, StreamExt}; +use tonic::Status; + +/// Response FIN extension. Absence preserves legacy response EOF behavior. +pub const HALF_CLOSE_CAPABILITY: &str = "stream-half-close-v1"; +/// `PeerRelay` response metadata confirms support along the downstream path. +pub const HALF_CLOSE_METADATA: &str = "openshell-stream-half-close"; +/// Outbound frames are bounded independently of peer frame-limit rollout. +pub const CHUNK_SIZE: usize = 64 * 1024; + +pub fn supports_half_close(capabilities: &[String]) -> bool { + capabilities + .iter() + .any(|value| value == HALF_CLOSE_CAPABILITY) +} + +pub fn capabilities() -> Vec { + vec![HALF_CLOSE_CAPABILITY.to_string()] +} + +/// Payloads after the first request frame. +pub enum Payload { + Data(Vec), + HalfClose, + Invalid, +} + +pub trait Frame: Send + 'static { + fn payload(self) -> Payload; + fn data(data: Vec) -> Self; + fn half_close() -> Self; +} + +macro_rules! frame { + ($name:ident, $module:ident) => { + impl Frame for crate::proto::$name { + fn payload(self) -> Payload { + match self.payload { + Some(crate::proto::$module::Payload::Data(data)) => Payload::Data(data), + Some(crate::proto::$module::Payload::HalfClose(_)) => Payload::HalfClose, + _ => Payload::Invalid, + } + } + fn data(data: Vec) -> Self { + Self { + payload: Some(crate::proto::$module::Payload::Data(data)), + } + } + fn half_close() -> Self { + Self { + payload: Some(crate::proto::$module::Payload::HalfClose( + crate::proto::StreamHalfClose {}, + )), + } + } + } + }; +} +frame!(TcpForwardFrame, tcp_forward_frame); +frame!(RelayFrame, relay_frame); +frame!(PeerRelayFrame, peer_relay_frame); + +pub fn close_status(close: &RelayClose) -> Status { + let code = match RelayCloseCode::try_from(close.code) { + Ok(RelayCloseCode::Cancelled) => tonic::Code::Cancelled, + Ok(RelayCloseCode::DeadlineExceeded) => tonic::Code::DeadlineExceeded, + Ok(RelayCloseCode::InvalidArgument) => tonic::Code::InvalidArgument, + _ => tonic::Code::Unavailable, + }; + Status::new(code, close.reason.clone()) +} + +pub fn close_message(channel_id: String, status: &Status) -> RelayClose { + let code = match status.code() { + tonic::Code::Cancelled => RelayCloseCode::Cancelled, + tonic::Code::DeadlineExceeded => RelayCloseCode::DeadlineExceeded, + tonic::Code::InvalidArgument => RelayCloseCode::InvalidArgument, + _ => RelayCloseCode::Unavailable, + }; + RelayClose { + channel_id, + reason: status.message().to_string(), + code: code as i32, + } +} + +#[derive(Default, Debug)] +struct AbortState { + error: Option, + waiters: Vec, +} + +/// First observed abort wins; waking both directions interrupts backpressure. +#[derive(Clone, Default, Debug)] +pub struct AbortHandle(Arc>); +impl AbortHandle { + pub fn abort(&self, error: Status) -> Status { + let waiters = { + let mut state = self.0.lock().unwrap(); + if let Some(first) = &state.error { + return first.clone(); + } + state.error = Some(error.clone()); + std::mem::take(&mut state.waiters) + }; + for waiter in waiters { + waiter.wake(); + } + error + } + fn poll_error(&self, cx: &Context<'_>) -> Option { + let mut state = self.0.lock().unwrap(); + if let Some(error) = &state.error { + return Some(error.clone()); + } + if !state.waiters.iter().any(|w| w.will_wake(cx.waker())) { + state.waiters.push(cx.waker().clone()); + } + None + } + pub async fn aborted(&self) -> Status { + std::future::poll_fn(|cx| self.poll_error(cx).map_or(Poll::Pending, Poll::Ready)).await + } + + /// Observe aborts even after socket EOF or while a frame send is blocked. + pub async fn run( + &self, + future: impl Future>, + ) -> Result<(), Status> { + let result = tokio::select! { + biased; + error = self.aborted() => Err(error), + result = future => result, + }; + result.map_err(|error| self.abort(error)) + } +} + +/// Internal byte pipe that never converts a known terminal failure into EOF. +#[derive(Debug)] +pub struct RelayIo { + stream: DuplexStream, + abort: AbortHandle, + half_close: bool, + completion: watch::Sender>>, +} + +/// Records RPC completion independently of directional byte EOF. Dropping an +/// unfinished producer must not leave an upstream caller waiting forever. +pub struct RelayCompletion { + result: watch::Sender>>, + abort: AbortHandle, +} +impl RelayCompletion { + pub fn finish(self, result: Result<(), Status>) { + let result = result.map_err(|error| self.abort.abort(error)); + self.result.send_replace(Some(result)); + } +} +impl Drop for RelayCompletion { + fn drop(&mut self) { + let error = Status::unavailable("relay ended before final status"); + let unfinished = self.result.send_if_modified(|result| { + if result.is_some() { + return false; + } + *result = Some(Err(error.clone())); + true + }); + if unfinished { + self.abort.abort(error); + } + } +} +impl RelayIo { + pub fn pair() -> (Self, Self) { + Self::pair_with_half_close(true) + } + + /// Carry downstream capability with the pipe, including legacy fallback. + pub fn pair_with_half_close(half_close: bool) -> (Self, Self) { + let (a, b) = tokio::io::duplex(CHUNK_SIZE); + let abort = AbortHandle::default(); + let (completion, _) = watch::channel(None); + ( + Self { + stream: a, + abort: abort.clone(), + half_close, + completion: completion.clone(), + }, + Self { + stream: b, + abort, + half_close, + completion, + }, + ) + } + pub fn abort_handle(&self) -> AbortHandle { + self.abort.clone() + } + + pub fn supports_half_close(&self) -> bool { + self.half_close + } + + /// The downstream bridge owns exactly one completion guard. + pub fn completion_guard(&self) -> RelayCompletion { + RelayCompletion { + result: self.completion.clone(), + abort: self.abort.clone(), + } + } + + /// Wait after forwarding FIN, before reporting successful RPC completion. + pub async fn completed(&self) -> Result<(), Status> { + let mut completion = self.completion.subscribe(); + loop { + let result = completion.borrow_and_update().clone(); + if let Some(result) = result { + return result; + } + completion + .changed() + .await + .map_err(|_| Status::unavailable("relay completion lost"))?; + } + } +} +impl AsyncRead for RelayIo { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + if let Some(error) = self.abort.poll_error(cx) { + return Poll::Ready(Err(io::Error::other(error))); + } + Pin::new(&mut self.stream).poll_read(cx, buf) + } +} +impl AsyncWrite for RelayIo { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + if let Some(error) = self.abort.poll_error(cx) { + return Poll::Ready(Err(io::Error::other(error))); + } + Pin::new(&mut self.stream).poll_write(cx, buf) + } + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if let Some(error) = self.abort.poll_error(cx) { + return Poll::Ready(Err(io::Error::other(error))); + } + Pin::new(&mut self.stream).poll_flush(cx) + } + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if let Some(error) = self.abort.poll_error(cx) { + return Poll::Ready(Err(io::Error::other(error))); + } + Pin::new(&mut self.stream).poll_shutdown(cx) + } +} + +pub fn io_status(error: io::Error) -> Status { + error + .get_ref() + .and_then(|e| e.downcast_ref::()) + .cloned() + .unwrap_or_else(|| Status::unavailable(error.to_string())) +} + +async fn receive_requests(mut inbound: S, mut write: W) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + W: AsyncWrite + Unpin, +{ + while let Some(frame) = inbound.next().await { + tokio::task::consume_budget().await; + let Payload::Data(data) = frame?.payload() else { + return Err(Status::invalid_argument( + "only data is allowed after init; close requests with EOF", + )); + }; + write.write_all(&data).await.map_err(io_status)?; + } + write.shutdown().await.map_err(io_status) +} + +async fn send_responses( + mut read: R, + tx: &mpsc::Sender>, + half_close: bool, +) -> Result<(), Status> { + let mut buf = vec![0; CHUNK_SIZE]; + loop { + let n = read.read(&mut buf).await.map_err(io_status)?; + if n == 0 { + break; + } + tx.send(Ok(F::data(buf[..n].to_vec()))) + .await + .map_err(|_| Status::cancelled("response dropped"))?; + } + if half_close { + tx.send(Ok(F::half_close())) + .await + .map_err(|_| Status::cancelled("response dropped"))?; + } + Ok(()) +} + +/// Reverse `RelayStream` carries target replies in the request direction. +/// +/// Legacy response EOF must reach the supervisor before its delayed reply can arrive. +/// Release the sole response sender at EOF, but keep draining requests. +pub async fn serve_relay( + inbound: S, + socket: T, + tx: &mut Option>>, + half_close: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + T: AsyncRead + AsyncWrite + Unpin, +{ + let sender = tx.as_ref().expect("relay response sender"); + if half_close { + return serve(inbound, socket, sender, true).await; + } + let (read, write) = tokio::io::split(socket); + let input = receive_requests(inbound, write); + tokio::pin!(input); + let input_finished = { + let output = send_responses(read, sender, false); + tokio::pin!(output); + let exchange = async { + tokio::select! { + result = &mut input => { result?; output.await?; Ok::<_, Status>(true) } + result = &mut output => { result?; Ok(false) } + } + }; + tokio::select! { + biased; + () = sender.closed() => return Err(Status::cancelled("response dropped")), + result = exchange => result?, + } + }; + tx.take(); + if !input_finished { + input.await?; + } + Ok(()) +} + +/// Forward-facing RPCs must retain downstream trailers after forwarding FIN. +pub async fn serve_forward( + inbound: S, + socket: &mut RelayIo, + tx: &mpsc::Sender>, + half_close: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, +{ + let abort = socket.abort_handle(); + abort + .run(async { + serve(inbound, &mut *socket, tx, half_close).await?; + // Legacy EOF may terminate a forward-facing RPC before request EOF. + // Waiting for its peer to drain those requests would deadlock. + if half_close { + tokio::select! { + biased; + () = tx.closed() => Err(Status::cancelled("response dropped")), + result = socket.completed() => result, + } + } else { + Ok(()) + } + }) + .await +} + +/// Run a server-side bridge after consuming init. Both futures stay owned by +/// this call, so dropping/cancelling it cannot leave a detached input pump. +pub async fn serve( + inbound: S, + socket: T, + tx: &mpsc::Sender>, + half_close: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + T: AsyncRead + AsyncWrite + Unpin, +{ + let (read, write) = tokio::io::split(socket); + let input = receive_requests(inbound, write); + let output = send_responses(read, tx, half_close); + tokio::pin!(input, output); + let exchange = async { + tokio::select! { + result = &mut input => { result?; output.await } + result = &mut output => { result?; if half_close { input.await } else { Ok(()) } } + } + }; + tokio::select! { + biased; + () = tx.closed() => Err(Status::cancelled("response dropped")), + result = exchange => result, + } +} + +/// Client-side bridge. Sending EOF does not cancel receiving. A negotiated +/// response FIN shuts down only the destination writer; trailers remain read. +pub async fn client( + inbound: S, + read: R, + write: W, + tx: mpsc::Sender, + half_close: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + client_impl(inbound, read, write, tx, half_close, false).await +} + +/// Internal relays retain their opposite pump on normal legacy response EOF. +/// In reverse `RelayStream` that pump carries the target's delayed reply. +pub async fn client_relay( + inbound: S, + read: R, + write: W, + tx: mpsc::Sender, + half_close: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + client_impl(inbound, read, write, tx, half_close, true).await +} + +async fn client_impl( + mut inbound: S, + mut read: R, + mut write: W, + tx: mpsc::Sender, + half_close: bool, + drain_input: bool, +) -> Result<(), Status> +where + F: Frame, + S: Stream> + Unpin, + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + let input = async move { + let mut buf = vec![0; CHUNK_SIZE]; + loop { + let n = read.read(&mut buf).await.map_err(io_status)?; + if n == 0 { + return Ok::<_, Status>(()); + } + tx.send(F::data(buf[..n].to_vec())) + .await + .map_err(|_| Status::unavailable("request stream closed"))?; + } + }; + let output = async { + let mut closed = false; + while let Some(frame) = inbound.next().await { + tokio::task::consume_budget().await; + match frame?.payload() { + Payload::Data(data) if !closed => { + write.write_all(&data).await.map_err(io_status)?; + write.flush().await.map_err(io_status)?; + } + Payload::HalfClose if half_close && !closed => { + write.shutdown().await.map_err(io_status)?; + closed = true; + } + _ => { + return Err(Status::invalid_argument( + "invalid response frame or data after half-close", + )); + } + } + } + if !closed { + write.shutdown().await.map_err(io_status)?; + } + Ok::<_, Status>(()) + }; + tokio::pin!(input, output); + tokio::select! { + biased; + result = &mut output => { result?; if drain_input { input.await } else { Ok(()) } }, + result = &mut input => { result?; output.await } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::proto::{PeerRelayFrame, RelayFrame, TcpForwardFrame}; + use std::time::Duration; + use tokio::net::{TcpListener, TcpStream}; + use tokio_stream::wrappers::ReceiverStream; + + async fn tcp_pair() -> (TcpStream, TcpStream) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dial = TcpStream::connect(listener.local_addr().unwrap()); + let (client, server) = tokio::join!(dial, listener.accept()); + (client.unwrap(), server.unwrap().0) + } + + async fn bounded(future: impl Future) -> T { + tokio::time::timeout(Duration::from_secs(5), future) + .await + .expect("bridge hung") + } + + #[tokio::test] + async fn request_fin_preserves_delayed_response() { + bounded(async { + let (bridge, mut target) = tcp_pair().await; + let (input, rx) = mpsc::channel(4); + let (output, mut responses) = mpsc::channel(4); + let task = tokio::spawn(async move { + serve::(ReceiverStream::new(rx), bridge, &output, true).await + }); + input.send(Ok(TcpForwardFrame::data(b"request".to_vec()))).await.unwrap(); + drop(input); + let mut request = Vec::new(); + target.read_to_end(&mut request).await.unwrap(); + assert_eq!(request, b"request"); + target.write_all(b"delayed response").await.unwrap(); + target.shutdown().await.unwrap(); + assert!(matches!(responses.recv().await.unwrap().unwrap().payload(), Payload::Data(data) if data == b"delayed response")); + assert!(matches!(responses.recv().await.unwrap().unwrap().payload(), Payload::HalfClose)); + assert!(responses.recv().await.is_none()); + task.await.unwrap().unwrap(); + }).await; + } + + #[tokio::test] + async fn response_fin_preserves_later_requests_across_three_hops() { + bounded(async { + let (client_socket, mut application) = tcp_pair().await; + let (supervisor_socket, mut target) = tcp_pair().await; + let (mut front, mut peer_client) = RelayIo::pair(); + let (mut peer_owner, mut relay_server) = RelayIo::pair(); + let (forward_tx, forward_rx) = mpsc::channel(4); + let (forward_out, forward_in) = mpsc::channel(4); + let (peer_tx, peer_rx) = mpsc::channel(4); + let (peer_out, peer_in) = mpsc::channel(4); + let (relay_tx, relay_rx) = mpsc::channel(4); + let (relay_out, relay_in) = mpsc::channel(4); + let (client_r, client_w) = tokio::io::split(client_socket); + let client_task = tokio::spawn(client::( + ReceiverStream::new(forward_in), + client_r, + client_w, + forward_tx, + true, + )); + let gateway = tokio::spawn(async move { + serve_forward::( + ReceiverStream::new(forward_rx).map(Ok), + &mut front, + &forward_out, + true, + ) + .await + }); + let peer = tokio::spawn(async move { + let completion = peer_client.completion_guard(); + let (r, w) = tokio::io::split(&mut peer_client); + let result = client_relay::( + ReceiverStream::new(peer_in), + r, + w, + peer_tx, + true, + ) + .await; + completion.finish(result.clone()); + result + }); + let owner = tokio::spawn(async move { + serve_forward::( + ReceiverStream::new(peer_rx).map(Ok), + &mut peer_owner, + &peer_out, + true, + ) + .await + }); + let relay = tokio::spawn(async move { + let completion = relay_server.completion_guard(); + let result = serve_relay::( + ReceiverStream::new(relay_rx).map(Ok), + &mut relay_server, + &mut Some(relay_out), + true, + ) + .await; + completion.finish(result.clone()); + result + }); + let (target_r, target_w) = tokio::io::split(supervisor_socket); + let supervisor = tokio::spawn(client_relay::( + ReceiverStream::new(relay_in), + target_r, + target_w, + relay_tx, + true, + )); + target.write_all(b"response").await.unwrap(); + target.shutdown().await.unwrap(); + let mut response = Vec::new(); + application.read_to_end(&mut response).await.unwrap(); + assert_eq!(response, b"response"); + application + .write_all(b"request after response FIN") + .await + .unwrap(); + application.shutdown().await.unwrap(); + let mut request = Vec::new(); + target.read_to_end(&mut request).await.unwrap(); + assert_eq!(request, b"request after response FIN"); + for result in [ + client_task.await, + gateway.await, + peer.await, + owner.await, + relay.await, + supervisor.await, + ] { + result.unwrap().unwrap(); + } + }) + .await; + } + + #[tokio::test] + async fn legacy_gateway_eof_preserves_supervisor_delayed_reply() { + bounded(async { + let (socket, mut target) = tcp_pair().await; + let (read, write) = tokio::io::split(socket); + let (requests, mut replies) = mpsc::channel(4); + // An old gateway ends RelayStream responses to signal input EOF. + let supervisor = tokio::spawn(client_relay::( + tokio_stream::iter([Ok(RelayFrame::data(b"request".to_vec()))]), + read, write, requests, false, + )); + let mut input = Vec::new(); + target.read_to_end(&mut input).await.unwrap(); + assert_eq!(input, b"request"); + assert!(!supervisor.is_finished()); + target.write_all(b"delayed reply").await.unwrap(); + target.shutdown().await.unwrap(); + assert!(matches!(replies.recv().await.unwrap().payload(), Payload::Data(data) if data == b"delayed reply")); + assert!(replies.recv().await.is_none()); + supervisor.await.unwrap().unwrap(); + }).await; + } + + #[tokio::test] + async fn legacy_supervisor_can_reply_after_gateway_response_eof() { + bounded(async { + let (mut caller, mut bridge) = RelayIo::pair_with_half_close(false); + let (requests, input) = mpsc::channel(4); + let (output, mut responses) = mpsc::channel(4); + let gateway = tokio::spawn(async move { + let completion = bridge.completion_guard(); + let result = serve_relay::(ReceiverStream::new(input), &mut bridge, &mut Some(output), false).await; + completion.finish(result.clone()); + result + }); + caller.write_all(b"request").await.unwrap(); + caller.shutdown().await.unwrap(); + assert!(matches!(responses.recv().await.unwrap().unwrap().payload(), Payload::Data(data) if data == b"request")); + assert!(responses.recv().await.is_none()); + // Old supervisors wait for response EOF before the target replies. + assert!(!gateway.is_finished()); + requests.send(Ok(RelayFrame::data(b"delayed reply".to_vec()))).await.unwrap(); + drop(requests); + let mut reply = Vec::new(); + caller.read_to_end(&mut reply).await.unwrap(); + assert_eq!(reply, b"delayed reply"); + caller.completed().await.unwrap(); + gateway.await.unwrap().unwrap(); + }).await; + } + + #[tokio::test] + async fn forward_waits_for_peer_trailers_after_fin() { + for terminal in [Ok(()), Err(Status::deadline_exceeded("late deadline"))] { + bounded(async { + let (mut front, mut bridge) = RelayIo::pair(); + let (output, mut responses) = mpsc::channel(4); + let gateway = tokio::spawn(async move { + serve_forward::( + tokio_stream::empty(), + &mut front, + &output, + true, + ) + .await + }); + let (peer_output, peer_responses) = mpsc::channel(4); + let (requests, mut peer_requests) = mpsc::channel(4); + let peer = tokio::spawn(async move { + let completion = bridge.completion_guard(); + let abort = bridge.abort_handle(); + let (read, write) = tokio::io::split(&mut bridge); + let result = abort + .run(client_relay::( + ReceiverStream::new(peer_responses), + read, + write, + requests, + true, + )) + .await; + completion.finish(result.clone()); + result + }); + assert!(peer_requests.recv().await.is_none()); + peer_output + .send(Ok(PeerRelayFrame::half_close())) + .await + .unwrap(); + assert!(matches!( + responses.recv().await.unwrap().unwrap().payload(), + Payload::HalfClose + )); + // Poll through the point where the old bridge returned success. + let mut gateway = gateway; + assert!( + tokio::time::timeout(Duration::from_millis(20), &mut gateway) + .await + .is_err() + ); + if let Err(error) = &terminal { + peer_output.send(Err(error.clone())).await.unwrap(); + } + drop(peer_output); + assert_eq!( + gateway.await.unwrap().map_err(|e| e.code()), + terminal.clone().map_err(|e| e.code()) + ); + assert_eq!( + peer.await.unwrap().map_err(|e| e.code()), + terminal.map_err(|e| e.code()) + ); + }) + .await; + } + } + + #[tokio::test] + async fn missing_final_status_does_not_hang_or_report_success() { + let (upstream, downstream) = RelayIo::pair(); + drop(downstream.completion_guard()); + assert_eq!( + bounded(upstream.abort_handle().aborted()).await.code(), + tonic::Code::Unavailable + ); + assert_eq!( + bounded(upstream.completed()).await.unwrap_err().code(), + tonic::Code::Unavailable + ); + } + + #[tokio::test] + async fn response_drop_cancels_wait_for_final_status() { + bounded(async { + let (mut upstream, mut downstream) = RelayIo::pair(); + let _completion = downstream.completion_guard(); + let abort = downstream.abort_handle(); + let (output, mut responses) = mpsc::channel(1); + let gateway = tokio::spawn(async move { + serve_forward::( + tokio_stream::empty(), + &mut upstream, + &output, + true, + ) + .await + }); + downstream.shutdown().await.unwrap(); + assert!(matches!( + responses.recv().await.unwrap().unwrap().payload(), + Payload::HalfClose + )); + drop(responses); + assert_eq!( + gateway.await.unwrap().unwrap_err().code(), + tonic::Code::Cancelled + ); + assert_eq!(abort.aborted().await.code(), tonic::Code::Cancelled); + }) + .await; + } + + #[tokio::test] + async fn legacy_downstream_drains_response_without_leaving_input_open() { + bounded(async { + let (bridge, mut downstream) = RelayIo::pair_with_half_close(false); + let half_close = bridge.supports_half_close(); + let (_input, rx) = mpsc::channel::>(1); + let (output, mut responses) = mpsc::channel(1); + let task = tokio::spawn(async move { + serve(ReceiverStream::new(rx), bridge, &output, half_close).await + }); + downstream.write_all(b"legacy response").await.unwrap(); + drop(downstream); + assert!(matches!(responses.recv().await.unwrap().unwrap().payload(), Payload::Data(data) if data == b"legacy response")); + // The caller advertised FIN, but downstream cannot preserve input. + // Drain bytes, then end the RPC without FIN or waiting for requests. + assert!(responses.recv().await.is_none()); + task.await.unwrap().unwrap(); + }).await; + } + + #[tokio::test] + async fn legacy_response_eof_finishes_without_request_eof() { + bounded(async { + let (bridge, mut target) = tcp_pair().await; + let (_input, rx) = mpsc::channel::>(1); + let (output, mut responses) = mpsc::channel(1); + let task = tokio::spawn(async move { + serve(ReceiverStream::new(rx), bridge, &output, false).await + }); + target.shutdown().await.unwrap(); + assert!(responses.recv().await.is_none()); + task.await.unwrap().unwrap(); + }) + .await; + } + + #[tokio::test] + async fn invalid_request_frames_are_protocol_errors() { + for frame in [ + TcpForwardFrame::default(), + TcpForwardFrame::half_close(), + TcpForwardFrame { + payload: Some(crate::proto::tcp_forward_frame::Payload::Init( + crate::proto::TcpForwardInit::default(), + )), + }, + ] { + let (socket, _target) = tokio::io::duplex(1); + let (out, _rx) = mpsc::channel(1); + let error = bounded(serve(tokio_stream::iter([Ok(frame)]), socket, &out, true)) + .await + .unwrap_err(); + assert_eq!(error.code(), tonic::Code::InvalidArgument); + } + } + + #[tokio::test] + async fn response_fin_does_not_hide_trailers_or_invalid_frames() { + for terminal in [ + Err(Status::deadline_exceeded("late deadline")), + Ok(TcpForwardFrame::data(vec![1])), + Ok(TcpForwardFrame::half_close()), + ] { + let expected = if terminal.is_err() { + tonic::Code::DeadlineExceeded + } else { + tonic::Code::InvalidArgument + }; + let (input, _rx) = mpsc::channel(1); + let (read, _peer) = tokio::io::duplex(1); + let error = bounded(client( + tokio_stream::iter([Ok(TcpForwardFrame::half_close()), terminal]), + read, + tokio::io::sink(), + input, + true, + )) + .await + .unwrap_err(); + assert_eq!(error.code(), expected); + } + } + + #[tokio::test] + async fn abort_wakes_blocked_read_and_write_and_preserves_first_status() { + bounded(async { + let (mut left, mut right) = RelayIo::pair(); + let abort = left.abort_handle(); + left.write_all(&vec![0; CHUNK_SIZE]).await.unwrap(); + let blocked_write = + tokio::spawn(async move { left.write_all(b"blocked").await.unwrap_err() }); + let blocked_read = tokio::spawn(async move { + right.write_all(&vec![0; CHUNK_SIZE + 1]).await.unwrap_err() + }); + abort.abort(Status::deadline_exceeded("deadline")); + abort.abort(Status::cancelled("later cancellation")); + assert_eq!( + io_status(blocked_write.await.unwrap()).code(), + tonic::Code::DeadlineExceeded + ); + assert_eq!( + io_status(blocked_read.await.unwrap()).code(), + tonic::Code::DeadlineExceeded + ); + let (mut a, _b) = RelayIo::pair(); + let abort = a.abort_handle(); + let read = tokio::spawn(async move { a.read(&mut [0]).await.unwrap_err() }); + abort.abort(Status::invalid_argument("invalid frame")); + assert_eq!( + io_status(read.await.unwrap()).code(), + tonic::Code::InvalidArgument + ); + }) + .await; + } + + #[tokio::test] + async fn dropping_response_cancels_a_backpressured_bridge() { + bounded(async { + let (bridge, mut target) = tokio::io::duplex(1); + let (input, rx) = mpsc::channel(1); + let (output, responses) = mpsc::channel(1); + let task = tokio::spawn(async move { + serve::(ReceiverStream::new(rx), bridge, &output, true).await + }); + input + .send(Ok(TcpForwardFrame::data(vec![0; 100]))) + .await + .unwrap(); + target.write_all(b"a").await.unwrap(); + drop(responses); + assert_eq!( + task.await.unwrap().unwrap_err().code(), + tonic::Code::Cancelled + ); + }) + .await; + } + + #[test] + fn unknown_capabilities_and_close_codes_have_legacy_fallback() { + assert!(!supports_half_close(&["future-feature".into()])); + assert!(supports_half_close(&capabilities())); + assert_eq!( + close_status(&RelayClose { + code: 999, + ..Default::default() + }) + .code(), + tonic::Code::Unavailable + ); + for code in [ + tonic::Code::Cancelled, + tonic::Code::DeadlineExceeded, + tonic::Code::InvalidArgument, + tonic::Code::Unavailable, + ] { + assert_eq!( + close_status(&close_message( + "channel".into(), + &Status::new(code, "reason") + )) + .code(), + code + ); + } + } +} diff --git a/crates/openshell-server/src/grpc/provider_readiness_tests.rs b/crates/openshell-server/src/grpc/provider_readiness_tests.rs index d83bb80ef4..9b89fd176c 100644 --- a/crates/openshell-server/src/grpc/provider_readiness_tests.rs +++ b/crates/openshell-server/src/grpc/provider_readiness_tests.rs @@ -72,6 +72,7 @@ fn hello() -> SupervisorHello { instance_id: Uuid::new_v4().to_string(), connection_epoch: 0, supports_provider_readiness: true, + ..Default::default() } } diff --git a/crates/openshell-server/src/grpc/sandbox.rs b/crates/openshell-server/src/grpc/sandbox.rs index 074b2590e0..1051cfdd8e 100644 --- a/crates/openshell-server/src/grpc/sandbox.rs +++ b/crates/openshell-server/src/grpc/sandbox.rs @@ -58,7 +58,7 @@ use tokio::net::{TcpListener, TcpStream}; use tokio::sync::{broadcast, mpsc, oneshot}; use tokio_stream::wrappers::ReceiverStream; use tonic::{Request, Response, Status}; -use tracing::{debug, info, warn}; +use tracing::{info, warn}; use russh::ChannelMsg; use russh::client::AuthResult; @@ -74,7 +74,6 @@ use super::validation::{ use super::{MAX_PROVIDERS, MAX_ROUTABLE_NAME_LEN}; use crate::persistence::current_time_ms; -const TCP_FORWARD_CHUNK_SIZE: usize = 64 * 1024; const NO_LOGIN_SHELL_ENV: (&str, &str) = ("OPENSHELL_NO_LOGIN_SHELL", "1"); const MAX_TEMPLATES_PER_WORKSPACE: u32 = 1000; const MAX_CREATE_SERVICE_EXPOSURES: usize = 32; @@ -2512,13 +2511,18 @@ pub(super) async fn handle_exec_sandbox( /// Returns `Some(stream)` on success. On any failure the error is sent on `tx` /// and `None` is returned; the caller should then `return` immediately. async fn await_relay_stream( - relay_rx: oneshot::Receiver>, + relay_rx: oneshot::Receiver>, tx: &mpsc::Sender>, sandbox_id: &str, channel_id: &str, context: &str, -) -> Option { - match tokio::time::timeout(std::time::Duration::from_secs(10), relay_rx).await { +) -> Option { + let result = tokio::select! { + biased; + () = tx.closed() => return None, + result = tokio::time::timeout(std::time::Duration::from_secs(10), relay_rx) => result, + }; + match result { Ok(Ok(Ok(stream))) => Some(stream), Ok(Ok(Err(status))) => { warn!(sandbox_id = %sandbox_id, channel_id = %channel_id, error = %status.message(), "{context}: relay target open failed"); @@ -2563,6 +2567,7 @@ pub(super) async fn handle_forward_tcp( )); }; + let half_close = openshell_core::stream_lifecycle::supports_half_close(&init.capabilities); let target = validate_tcp_forward_init(&init)?; let sandbox = resolve_and_authorize_sandbox_name( @@ -2602,7 +2607,7 @@ pub(super) async fn handle_forward_tcp( return; }; - bridge_forward_tcp_stream(inbound, relay_stream, tx, &sandbox_id, &channel_id).await; + bridge_forward_tcp_stream(inbound, relay_stream, tx, half_close).await; }); let stream: Pin< @@ -2777,70 +2782,24 @@ fn validate_tcp_target_parts(host: &str, _port: u32) -> Result { } async fn bridge_forward_tcp_stream( - mut inbound: tonic::Streaming, - relay_stream: tokio::io::DuplexStream, + inbound: tonic::Streaming, + mut relay_stream: openshell_core::stream_lifecycle::RelayIo, tx: mpsc::Sender>, - sandbox_id: &str, - channel_id: &str, + half_close: bool, ) { - let (mut relay_read, mut relay_write) = tokio::io::split(relay_stream); - - let sandbox_id_in = sandbox_id.to_string(); - let channel_id_in = channel_id.to_string(); - tokio::spawn(async move { - loop { - match inbound.message().await { - Ok(Some(frame)) => { - let Some(openshell_core::proto::tcp_forward_frame::Payload::Data(data)) = - frame.payload - else { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, "ForwardTcp: received non-data frame after init"); - break; - }; - if data.is_empty() { - continue; - } - if let Err(err) = - tokio::io::AsyncWriteExt::write_all(&mut relay_write, &data).await - { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, error = %err, "ForwardTcp: write to relay failed"); - break; - } - } - Ok(None) => break, - Err(err) => { - debug!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, error = %err, "ForwardTcp: inbound stream ended"); - break; - } - } - } - let _ = tokio::io::AsyncWriteExt::shutdown(&mut relay_write).await; - }); - - let mut buf = vec![0u8; TCP_FORWARD_CHUNK_SIZE]; - loop { - match tokio::io::AsyncReadExt::read(&mut relay_read, &mut buf).await { - Ok(0) => break, - Ok(n) => { - let frame = TcpForwardFrame { - payload: Some(openshell_core::proto::tcp_forward_frame::Payload::Data( - buf[..n].to_vec(), - )), - }; - if tx.send(Ok(frame)).await.is_err() { - break; - } - } - Err(err) => { - warn!(sandbox_id = %sandbox_id, channel_id = %channel_id, error = %err, "ForwardTcp: read from relay failed"); - let _ = tx - .send(Err(Status::unavailable(format!( - "relay read failed: {err}" - )))) - .await; - break; - } - } + let half_close = half_close && relay_stream.supports_half_close(); + let abort = relay_stream.abort_handle(); + let result = abort + .run(openshell_core::stream_lifecycle::serve_forward( + inbound, + &mut relay_stream, + &tx, + half_close, + )) + .await; + if let Err(status) = result { + abort.abort(status.clone()); + let _ = tx.send(Err(status)).await; } } @@ -3314,7 +3273,7 @@ async fn stream_exec_over_relay( tx: mpsc::Sender>, sandbox_id: &str, channel_id: &str, - relay_stream: tokio::io::DuplexStream, + relay_stream: openshell_core::stream_lifecycle::RelayIo, command: &str, stdin_payload: Vec, execution_timeout: Option, @@ -3392,7 +3351,7 @@ async fn stream_interactive_exec_over_relay( tx: mpsc::Sender>, sandbox_id: &str, channel_id: &str, - relay_stream: tokio::io::DuplexStream, + relay_stream: openshell_core::stream_lifecycle::RelayIo, command: &str, input_stream: tonic::Streaming, request_tty: bool, @@ -3716,7 +3675,7 @@ async fn run_interactive_exec_with_russh( /// The supervisor bridges the relay to its Unix-socket SSH daemon; filesystem /// permissions on that socket are the only access-control boundary. async fn start_single_use_ssh_proxy_over_relay( - mut relay_stream: tokio::io::DuplexStream, + mut relay_stream: openshell_core::stream_lifecycle::RelayIo, ) -> Result<(u16, tokio::task::JoinHandle<()>), Box> { let listener = TcpListener::bind(("127.0.0.1", 0)).await?; let port = listener.local_addr()?.port(); @@ -4197,6 +4156,7 @@ mod tests { port: 8080, })), authorization_token: String::new(), + capabilities: Vec::new(), }; validate_tcp_forward_init(&init).expect("loopback target should pass"); } @@ -4227,6 +4187,7 @@ mod tests { port: 8080, })), authorization_token: String::new(), + capabilities: Vec::new(), }; assert_eq!( validate_tcp_forward_init(&init) @@ -4247,6 +4208,7 @@ mod tests { port: 0, })), authorization_token: String::new(), + capabilities: Vec::new(), }; assert_eq!( validate_tcp_forward_init(&init) diff --git a/crates/openshell-server/src/supervisor_session.rs b/crates/openshell-server/src/supervisor_session.rs index 710ecc1406..b517e06719 100644 --- a/crates/openshell-server/src/supervisor_session.rs +++ b/crates/openshell-server/src/supervisor_session.rs @@ -17,7 +17,7 @@ use uuid::Uuid; use openshell_core::proto::{ GatewayMessage, GetSandboxProviderStatusRequest, GetSandboxProviderStatusResponse, - PeerRelayFrame, PeerRelayInit, ProviderReadinessObservation, RelayFrame, RelayInit, RelayOpen, + PeerRelayFrame, PeerRelayInit, ProviderReadinessObservation, RelayFrame, RelayOpen, ReportEndpointStatusRequest, ReportEndpointStatusResponse, ReportMainProcessExitRequest, ReportMainProcessExitResponse, ReportProviderReadinessRequest, ReportProviderReadinessResponse, Sandbox, SandboxPhase, SessionAccepted, SshRelayTarget, SupervisorHello, SupervisorMessage, @@ -30,6 +30,7 @@ use crate::auth::principal::Principal; use crate::grpc::provider_readiness::ProviderReadinessEvidence; use crate::persistence::ObjectId; use crate::supervisor_owner::{OWNER_TTL, OwnerError, OwnerGuard, SupervisorOwnerIndex}; +use openshell_core::stream_lifecycle::{self, AbortHandle, RelayIo}; const HEARTBEAT_INTERVAL_SECS: u32 = 15; const OWNER_RENEW_TIMEOUT: Duration = Duration::from_secs(5); @@ -315,7 +316,7 @@ pub(crate) struct EndpointReportCursor { /// Holds a oneshot sender that will deliver the upgraded relay stream or a /// target-open failure reported by the supervisor. -type RelayStreamSender = oneshot::Sender>; +type RelayStreamSender = oneshot::Sender>; /// Registry of active supervisor sessions and pending relay channels. #[derive(Default)] @@ -329,6 +330,25 @@ pub struct SupervisorSessionRegistry { session_lifetimes: Arc>, admission_closed: AtomicBool, shutdown: watch::Sender, + active_relays: Arc>>, +} + +#[derive(Debug)] +struct ActiveRelay { + sandbox_id: String, + session_id: String, + abort: AbortHandle, +} + +#[derive(Debug)] +pub struct ActiveRelayGuard { + channels: Arc>>, + channel_id: String, +} +impl Drop for ActiveRelayGuard { + fn drop(&mut self) { + self.channels.lock().unwrap().remove(&self.channel_id); + } } struct PendingRelay { @@ -340,8 +360,10 @@ struct PendingRelay { #[derive(Debug)] pub struct ClaimedRelay { - pub stream: tokio::io::DuplexStream, + pub guard: ActiveRelayGuard, + pub stream: RelayIo, pub sandbox_id: String, + session_id: String, } impl std::fmt::Debug for SupervisorSessionRegistry { @@ -737,13 +759,7 @@ impl SupervisorSessionRegistry { &self, sandbox_id: &str, session_wait_timeout: Duration, - ) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, - > { + ) -> Result<(String, oneshot::Receiver>), Status> { self.open_relay_with_target( sandbox_id, relay_open::Target::Ssh(SshRelayTarget {}), @@ -759,13 +775,7 @@ impl SupervisorSessionRegistry { target: relay_open::Target, service_id: String, session_wait_timeout: Duration, - ) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, - > { + ) -> Result<(String, oneshot::Receiver>), Status> { let channel_id = Uuid::new_v4().to_string(); let relay_open = RelayOpen { channel_id: channel_id.clone(), @@ -781,13 +791,7 @@ impl SupervisorSessionRegistry { sandbox_id: &str, relay_open: RelayOpen, session_wait_timeout: Duration, - ) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, - > { + ) -> Result<(String, oneshot::Receiver>), Status> { if relay_open.channel_id.is_empty() { return Err(Status::invalid_argument("relay channel_id is required")); } @@ -803,7 +807,15 @@ impl SupervisorSessionRegistry { // both insert past it. let (relay_tx, relay_rx) = oneshot::channel(); { + // Serialize allocation with claim activation. PeerRelay supplies + // its channel ID, so it must not replace a pending or active relay. + let _sessions = self.sessions.lock().unwrap(); let mut pending = self.pending_relays.lock().unwrap(); + if pending.contains_key(&channel_id) + || self.active_relays.lock().unwrap().contains_key(&channel_id) + { + return Err(Status::already_exists("relay channel already exists")); + } if pending.len() >= MAX_PENDING_RELAYS { return Err(Status::resource_exhausted(format!( "gateway relay capacity reached ({MAX_PENDING_RELAYS} in flight)" @@ -842,6 +854,47 @@ impl SupervisorSessionRegistry { Ok((channel_id, relay_rx)) } + fn abort_relay( + &self, + sandbox_id: &str, + session_id: &str, + close: &openshell_core::proto::RelayClose, + ) -> bool { + let sessions = self.sessions.lock().unwrap(); + if sessions + .get(sandbox_id) + .is_none_or(|session| session.session_id != session_id) + { + return false; + } + // Keep the session lock through cancellation so a concurrent claim + // cannot move the channel from pending to active between these checks. + { + let mut pending = self.pending_relays.lock().unwrap(); + if let Some(relay) = pending.get(&close.channel_id) { + if relay.sandbox_id != sandbox_id { + return false; + } + let relay = pending + .remove(&close.channel_id) + .expect("pending relay existed before removal"); + let _ = relay + .sender + .send(Err(stream_lifecycle::close_status(close))); + return true; + } + } + let active = self.active_relays.lock().unwrap(); + let Some(relay) = active + .get(&close.channel_id) + .filter(|relay| relay.sandbox_id == sandbox_id && relay.session_id == session_id) + else { + return false; + }; + relay.abort.abort(stream_lifecycle::close_status(close)); + true + } + pub fn fail_pending_relay(&self, channel_id: &str, error: String) -> bool { let pending = self.pending_relays.lock().unwrap().remove(channel_id); if let Some(pending) = pending { @@ -863,6 +916,17 @@ impl SupervisorSessionRegistry { channel_id: &str, principal: Option<&Principal>, ) -> Result { + self.claim_relay_for_session(channel_id, principal, "", false) + } + + fn claim_relay_for_session( + &self, + channel_id: &str, + principal: Option<&Principal>, + session_id: &str, + half_close: bool, + ) -> Result { + let sessions = self.sessions.lock().unwrap(); let pending = { let mut map = self.pending_relays.lock().unwrap(); let pending = map @@ -883,6 +947,16 @@ impl SupervisorSessionRegistry { return Err(status); } + if !session_id.is_empty() + && sessions + .get(&pending.sandbox_id) + .is_none_or(|session| session.session_id != session_id) + { + return Err(Status::failed_precondition( + "relay belongs to a stale supervisor session", + )); + } + if pending.created_at.elapsed() > RELAY_PENDING_TIMEOUT { map.remove(channel_id); return Err(Status::deadline_exceeded("relay channel timed out")); @@ -894,16 +968,34 @@ impl SupervisorSessionRegistry { // Create a duplex stream pair: one end for the gateway bridge, one for // the supervisor HTTP CONNECT handler. - let (gateway_stream, supervisor_stream) = tokio::io::duplex(64 * 1024); + let (gateway_stream, supervisor_stream) = RelayIo::pair_with_half_close(half_close); + let session_id = sessions + .get(&pending.sandbox_id) + .map(|session| session.session_id.clone()) + .unwrap_or_default(); + self.active_relays.lock().unwrap().insert( + channel_id.to_string(), + ActiveRelay { + sandbox_id: pending.sandbox_id.clone(), + session_id: session_id.clone(), + abort: gateway_stream.abort_handle(), + }, + ); + let guard = ActiveRelayGuard { + channels: Arc::clone(&self.active_relays), + channel_id: channel_id.to_string(), + }; // Send the gateway-side stream to the waiter (exec handler or forward handler). if pending.sender.send(Ok(gateway_stream)).is_err() { return Err(Status::internal("relay requester dropped")); } Ok(ClaimedRelay { + guard, stream: supervisor_stream, sandbox_id: pending.sandbox_id, + session_id, }) } @@ -989,10 +1081,7 @@ async fn require_persisted_sandbox( // RelayStream gRPC handler // --------------------------------------------------------------------------- -/// Size of chunks read from the gateway-side `DuplexStream` when forwarding -/// bytes back to the supervisor over the gRPC response stream. -const RELAY_STREAM_CHUNK_SIZE: usize = 16 * 1024; - +/// Response frames and terminal status for a supervisor relay. type RelayStreamResponse = Response< Pin> + Send + 'static>>, >; @@ -1025,114 +1114,82 @@ async fn handle_relay_stream_inner( let principal = request.extensions().get::().cloned(); let mut inbound = request.into_inner(); - // First frame must identify the channel. let first = inbound .message() .await? .ok_or_else(|| Status::invalid_argument("empty RelayStream"))?; - let channel_id = match first.payload { - Some(openshell_core::proto::relay_frame::Payload::Init(RelayInit { channel_id })) - if !channel_id.is_empty() => - { - channel_id - } - _ => { - return Err(Status::invalid_argument( - "first RelayFrame must be init with non-empty channel_id", - )); - } + let Some(openshell_core::proto::relay_frame::Payload::Init(init)) = first.payload else { + return Err(Status::invalid_argument("first RelayFrame must be init")); }; - - // Claim the pending relay. Consumes the entry — it cannot be reused. - let claimed = registry.claim_relay(&channel_id, principal.as_ref())?; - let sandbox_id = claimed.sandbox_id; - let supervisor_side = claimed.stream; - info!(channel_id = %channel_id, sandbox_id = %sandbox_id, "relay stream: claimed pending relay, bridging"); - - let (mut read_half, mut write_half) = tokio::io::split(supervisor_side); - - // Supervisor → gateway: drain `inbound` and write to the DuplexStream. - let channel_id_in = channel_id.clone(); - let sandbox_id_in = sandbox_id; - let state_in = state.clone(); + if init.channel_id.is_empty() { + return Err(Status::invalid_argument("channel_id is required")); + } + let half_close = stream_lifecycle::supports_half_close(&init.capabilities); + if half_close && init.session_id.is_empty() { + return Err(Status::invalid_argument( + "negotiated relay requires session_id", + )); + } + let mut claimed = registry.claim_relay_for_session( + &init.channel_id, + principal.as_ref(), + &init.session_id, + half_close, + )?; + let (out_tx, out_rx) = mpsc::channel(16); + let completion = claimed.stream.completion_guard(); tokio::spawn(async move { - loop { - match inbound.message().await { - Ok(Some(frame)) => { - let Some(openshell_core::proto::relay_frame::Payload::Data(data)) = - frame.payload - else { - warn!(channel_id = %channel_id_in, "relay stream: received non-data frame after init"); - break; - }; - if data.is_empty() { - continue; - } - if let Err(e) = - tokio::io::AsyncWriteExt::write_all(&mut write_half, &data).await - { - warn!(channel_id = %channel_id_in, error = %e, "relay stream: write to duplex failed"); - break; - } - } - Ok(None) => break, - Err(e) => { - if let Some(state) = state_in.as_ref() - && expected_transport_close_during_sandbox_teardown( - state, - &sandbox_id_in, - &e, - ) - .await - { - info!( - sandbox_id = %sandbox_id_in, - channel_id = %channel_id_in, - error = %e, - "relay stream: expected transport close during sandbox teardown" - ); - } else { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, error = %e, "relay stream: inbound errored"); - } - break; - } + let mut out_tx = Some(out_tx); + let _guard = claimed.guard; + let abort = claimed.stream.abort_handle(); + let result = abort + .run(stream_lifecycle::serve_relay( + inbound, + &mut claimed.stream, + &mut out_tx, + half_close, + )) + .await; + completion.finish(result.clone()); + if let Err(status) = result { + abort.abort(status.clone()); + let expected_close = if let Some(state) = &state { + expected_transport_close_during_sandbox_teardown( + state, + &claimed.sandbox_id, + &status, + ) + .await + } else { + false + }; + if expected_close { + debug!(channel_id = %init.channel_id, "relay closed during sandbox teardown"); } - } - // Best-effort half-close on the write side so the reader sees EOF. - let _ = tokio::io::AsyncWriteExt::shutdown(&mut write_half).await; - }); - - // Gateway → supervisor: read the DuplexStream and emit RelayFrame::data messages. - let (out_tx, out_rx) = mpsc::channel::>(16); - let channel_id_out = channel_id; - tokio::spawn(async move { - let mut buf = vec![0u8; RELAY_STREAM_CHUNK_SIZE]; - loop { - match tokio::io::AsyncReadExt::read(&mut read_half, &mut buf).await { - Ok(0) => break, - Ok(n) => { - let chunk = RelayFrame { - payload: Some(openshell_core::proto::relay_frame::Payload::Data( - buf[..n].to_vec(), + // Typed abort is observable through both the byte pipe and trailers. + if let Some(state) = state { + let tx = state + .supervisor_sessions + .sessions + .lock() + .unwrap() + .get(&claimed.sandbox_id) + .filter(|session| session.session_id == claimed.session_id) + .map(|session| session.tx.clone()); + if let Some(tx) = tx { + let _ = tx.try_send(GatewayMessage { + payload: Some(gateway_message::Payload::RelayClose( + stream_lifecycle::close_message(init.channel_id, &status), )), - }; - if out_tx.send(Ok(chunk)).await.is_err() { - break; - } - } - Err(e) => { - warn!(channel_id = %channel_id_out, error = %e, "relay stream: read from duplex failed"); - break; + }); } } + if let Some(out_tx) = out_tx { + let _ = out_tx.send(Err(status)).await; + } } }); - - let stream = ReceiverStream::new(out_rx); - let stream: Pin< - Box> + Send + 'static>, - > = Box::pin(stream); - Ok(Response::new(stream)) + Ok(Response::new(Box::pin(ReceiverStream::new(out_rx)))) } fn expected_transport_close_during_shutdown(status: &Status, terminating: bool) -> bool { @@ -1362,13 +1419,7 @@ pub async fn open_routed_relay_with_target( target: relay_open::Target, service_id: String, session_wait_timeout: Duration, -) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, -> { +) -> Result<(String, oneshot::Receiver>), Status> { let channel_id = Uuid::new_v4().to_string(); let relay_open = RelayOpen { channel_id: channel_id.clone(), @@ -1383,13 +1434,7 @@ pub async fn open_routed_relay_with_message( sandbox_id: &str, relay_open: RelayOpen, session_wait_timeout: Duration, -) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, -> { +) -> Result<(String, oneshot::Receiver>), Status> { let deadline = Instant::now() + session_wait_timeout; let mut backoff = SESSION_WAIT_INITIAL_BACKOFF; let owner_index = SupervisorOwnerIndex::new(state.store.clone(), OWNER_TTL); @@ -1510,13 +1555,7 @@ async fn open_peer_relay( owner_peer_endpoint: String, sandbox_id: &str, relay_open: RelayOpen, -) -> Result< - ( - String, - oneshot::Receiver>, - ), - Status, -> { +) -> Result<(String, oneshot::Receiver>), Status> { let channel_id = relay_open.channel_id.clone(); let (relay_tx, relay_rx) = oneshot::channel(); let stream = connect_peer_relay(state, &owner_peer_endpoint, sandbox_id, relay_open).await?; @@ -1529,7 +1568,7 @@ async fn connect_peer_relay( owner_peer_endpoint: &str, sandbox_id: &str, relay_open: RelayOpen, -) -> Result { +) -> Result { let token = state.peer_routes.peer_token().await?; let channel = state.peer_routes.channel(owner_peer_endpoint).await?; let interceptor = PeerAuthInterceptor::new(&token, &state.replica_id)?; @@ -1542,6 +1581,7 @@ async fn connect_peer_relay( sandbox_id: sandbox_id.to_string(), relay_open: Some(relay_open), requester_replica_id: state.replica_id.clone(), + capabilities: stream_lifecycle::capabilities(), })), }) .await @@ -1554,8 +1594,12 @@ async fn connect_peer_relay( state.peer_routes.evict_channel(owner_peer_endpoint); Status::unavailable(format!("gateway peer relay RPC failed: {err}")) })?; + let half_close = response + .metadata() + .get(stream_lifecycle::HALF_CLOSE_METADATA) + .is_some_and(|value| value == "v1"); let inbound = response.into_inner(); - let (gateway_stream, bridge_stream) = tokio::io::duplex(64 * 1024); + let (gateway_stream, bridge_stream) = RelayIo::pair_with_half_close(half_close); spawn_peer_bridge(bridge_stream, inbound, out_tx, sandbox_id.to_string()); Ok(gateway_stream) } @@ -1596,6 +1640,7 @@ pub async fn handle_peer_relay( "peer relay requester does not match authenticated gateway replica", )); } + let half_close = stream_lifecycle::supports_half_close(&init.capabilities); let relay_open = init .relay_open .ok_or_else(|| Status::invalid_argument("relay_open is required"))?; @@ -1621,6 +1666,7 @@ pub async fn handle_peer_relay( Err(_) => return Err(Status::deadline_exceeded("relay open timed out")), }; + let half_close = half_close && supervisor_stream.supports_half_close(); let (out_tx, out_rx) = mpsc::channel::>(16); spawn_peer_owner_bridge( supervisor_stream, @@ -1628,133 +1674,62 @@ pub async fn handle_peer_relay( out_tx, init.sandbox_id, channel_id, + half_close, ); let stream: Pin< Box> + Send + 'static>, > = Box::pin(ReceiverStream::new(out_rx)); - Ok(Response::new(stream)) + let mut response = Response::new(stream); + if half_close { + response.metadata_mut().insert( + stream_lifecycle::HALF_CLOSE_METADATA, + MetadataValue::from_static("v1"), + ); + } + Ok(response) } fn spawn_peer_bridge( - bridge_stream: tokio::io::DuplexStream, - mut inbound: tonic::Streaming, + mut bridge_stream: RelayIo, + inbound: tonic::Streaming, out_tx: mpsc::Sender, - sandbox_id: String, + _sandbox_id: String, ) { - let (mut read_half, mut write_half) = tokio::io::split(bridge_stream); - let sandbox_id_in = sandbox_id.clone(); + let completion = bridge_stream.completion_guard(); tokio::spawn(async move { - loop { - match inbound.message().await { - Ok(Some(frame)) => { - let Some(peer_relay_frame::Payload::Data(data)) = frame.payload else { - warn!(sandbox_id = %sandbox_id_in, "gateway peer relay: non-data frame after init"); - break; - }; - if data.is_empty() { - continue; - } - if let Err(err) = - tokio::io::AsyncWriteExt::write_all(&mut write_half, &data).await - { - warn!(sandbox_id = %sandbox_id_in, error = %err, "gateway peer relay: write to duplex failed"); - break; - } - } - Ok(None) => break, - Err(err) => { - warn!(sandbox_id = %sandbox_id_in, error = %err, "gateway peer relay: inbound errored"); - break; - } - } - } - let _ = tokio::io::AsyncWriteExt::shutdown(&mut write_half).await; - }); - - tokio::spawn(async move { - let mut buf = vec![0u8; RELAY_STREAM_CHUNK_SIZE]; - loop { - match tokio::io::AsyncReadExt::read(&mut read_half, &mut buf).await { - Ok(0) => break, - Ok(n) => { - if out_tx - .send(PeerRelayFrame { - payload: Some(peer_relay_frame::Payload::Data(buf[..n].to_vec())), - }) - .await - .is_err() - { - break; - } - } - Err(err) => { - warn!(sandbox_id = %sandbox_id, error = %err, "gateway peer relay: read from duplex failed"); - break; - } - } - } + let abort = bridge_stream.abort_handle(); + let half_close = bridge_stream.supports_half_close(); + let (read, write) = tokio::io::split(&mut bridge_stream); + let result = abort + .run(stream_lifecycle::client_relay( + inbound, read, write, out_tx, half_close, + )) + .await; + completion.finish(result); }); } fn spawn_peer_owner_bridge( - supervisor_stream: tokio::io::DuplexStream, - mut inbound: tonic::Streaming, + mut supervisor_stream: RelayIo, + inbound: tonic::Streaming, out_tx: mpsc::Sender>, - sandbox_id: String, - channel_id: String, + _sandbox_id: String, + _channel_id: String, + half_close: bool, ) { - let (mut read_half, mut write_half) = tokio::io::split(supervisor_stream); - let sandbox_id_in = sandbox_id.clone(); - let channel_id_in = channel_id.clone(); - tokio::spawn(async move { - loop { - match inbound.message().await { - Ok(Some(frame)) => { - let Some(peer_relay_frame::Payload::Data(data)) = frame.payload else { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, "gateway peer relay owner: non-data frame after init"); - break; - }; - if data.is_empty() { - continue; - } - if let Err(err) = - tokio::io::AsyncWriteExt::write_all(&mut write_half, &data).await - { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, error = %err, "gateway peer relay owner: write to supervisor relay failed"); - break; - } - } - Ok(None) => break, - Err(err) => { - warn!(sandbox_id = %sandbox_id_in, channel_id = %channel_id_in, error = %err, "gateway peer relay owner: inbound errored"); - break; - } - } - } - let _ = tokio::io::AsyncWriteExt::shutdown(&mut write_half).await; - }); - tokio::spawn(async move { - let mut buf = vec![0u8; RELAY_STREAM_CHUNK_SIZE]; - loop { - match tokio::io::AsyncReadExt::read(&mut read_half, &mut buf).await { - Ok(0) => break, - Ok(n) => { - if out_tx - .send(Ok(PeerRelayFrame { - payload: Some(peer_relay_frame::Payload::Data(buf[..n].to_vec())), - })) - .await - .is_err() - { - break; - } - } - Err(err) => { - warn!(sandbox_id = %sandbox_id, channel_id = %channel_id, error = %err, "gateway peer relay owner: read from supervisor relay failed"); - break; - } - } + let abort = supervisor_stream.abort_handle(); + if let Err(error) = abort + .run(stream_lifecycle::serve_forward( + inbound, + &mut supervisor_stream, + &out_tx, + half_close, + )) + .await + { + abort.abort(error.clone()); + let _ = out_tx.send(Err(error)).await; } }); } @@ -1914,6 +1889,11 @@ async fn establish_supervisor_session( let accepted = GatewayMessage { payload: Some(gateway_message::Payload::SessionAccepted(SessionAccepted { session_id: session_id.clone(), + capabilities: if stream_lifecycle::supports_half_close(&hello.capabilities) { + stream_lifecycle::capabilities() + } else { + Vec::new() + }, heartbeat_interval: openshell_core::time::duration_from_std(Duration::from_secs( u64::from(HEARTBEAT_INTERVAL_SECS), )) @@ -2241,6 +2221,9 @@ async fn handle_supervisor_message( } } Some(supervisor_message::Payload::RelayClose(close)) => { + state + .supervisor_sessions + .abort_relay(sandbox_id, session_id, &close); info!( sandbox_id = %sandbox_id, session_id = %session_id, @@ -3045,6 +3028,163 @@ mod tests { assert!(!sandbox_proto_is_terminating(&sandbox)); } + #[tokio::test] + async fn typed_close_requires_the_owning_sandbox_and_session() { + let registry = SupervisorSessionRegistry::new(); + let (tx, _rx) = mpsc::channel(4); + registry.register( + "sbx-test".into(), + "current".into(), + tx.clone(), + make_shutdown(), + ); + registry.register( + "other".into(), + "other-session".into(), + tx.clone(), + make_shutdown(), + ); + let (relay_tx, relay_rx) = oneshot::channel(); + registry.pending_relays.lock().unwrap().insert( + "ch".into(), + pending_relay("sbx-test", relay_tx, Instant::now()), + ); + let principal = sandbox_principal("sbx-test"); + assert_eq!( + registry + .claim_relay_for_session("ch", Some(&principal), "stale", true) + .unwrap_err() + .code(), + tonic::Code::FailedPrecondition + ); + assert!(registry.pending_relays.lock().unwrap().contains_key("ch")); + let claimed = registry + .claim_relay_for_session("ch", Some(&principal), "current", true) + .unwrap(); + let mut pipe = relay_rx.await.unwrap().unwrap(); + let close = + stream_lifecycle::close_message("ch".into(), &Status::deadline_exceeded("deadline")); + assert!(!registry.abort_relay("other", "other-session", &close)); + assert!(!registry.abort_relay("sbx-test", "stale", &close)); + assert!(registry.abort_relay("sbx-test", "current", &close)); + let error = pipe.read(&mut [0]).await.unwrap_err(); + assert_eq!( + stream_lifecycle::io_status(error).code(), + tonic::Code::DeadlineExceeded + ); + registry.register("sbx-test".into(), "replacement".into(), tx, make_shutdown()); + assert!(!registry.abort_relay("sbx-test", "current", &close)); + assert!(!registry.abort_relay("sbx-test", "replacement", &close)); + drop(claimed); + assert!(registry.active_relays.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn typed_close_before_claim_reports_error_and_releases_capacity() { + let registry = SupervisorSessionRegistry::new(); + let (tx, _rx) = mpsc::channel(4); + registry.register("sbx-test".into(), "current".into(), tx, make_shutdown()); + { + let mut pending = registry.pending_relays.lock().unwrap(); + for i in 0..MAX_PENDING_RELAYS_PER_SANDBOX - 1 { + let (sender, _) = oneshot::channel(); + pending.insert( + format!("channel-{i}"), + pending_relay("sbx-test", sender, Instant::now()), + ); + } + } + let (channel_id, mut relay_rx) = registry + .open_relay("sbx-test", Duration::from_secs(1)) + .await + .unwrap(); + assert_eq!( + registry + .open_relay("sbx-test", Duration::from_secs(1)) + .await + .unwrap_err() + .code(), + tonic::Code::ResourceExhausted + ); + + let close = stream_lifecycle::close_message( + channel_id.clone(), + &Status::deadline_exceeded("relay setup failed"), + ); + assert!(registry.abort_relay("sbx-test", "current", &close)); + let error = relay_rx.try_recv().unwrap().unwrap_err(); + assert_eq!(error.code(), tonic::Code::DeadlineExceeded); + assert_eq!(error.message(), "relay setup failed"); + assert!( + !registry + .pending_relays + .lock() + .unwrap() + .contains_key(&channel_id) + ); + assert_eq!( + registry + .claim_relay_for_session( + &channel_id, + Some(&sandbox_principal("sbx-test")), + "current", + true, + ) + .unwrap_err() + .code(), + tonic::Code::NotFound + ); + assert!(!registry.abort_relay("sbx-test", "current", &close)); + registry + .open_relay("sbx-test", Duration::from_secs(1)) + .await + .expect("closing a pending relay should immediately free capacity"); + } + + #[tokio::test] + async fn typed_close_before_claim_requires_current_session_and_matching_sandbox() { + let registry = SupervisorSessionRegistry::new(); + let (tx, _rx) = mpsc::channel(4); + registry.register("sbx-test".into(), "old".into(), tx.clone(), make_shutdown()); + let (channel_id, mut relay_rx) = registry + .open_relay("sbx-test", Duration::from_secs(1)) + .await + .unwrap(); + registry.register( + "sbx-test".into(), + "current".into(), + tx.clone(), + make_shutdown(), + ); + registry.register("other".into(), "other-session".into(), tx, make_shutdown()); + let close = stream_lifecycle::close_message( + channel_id.clone(), + &Status::cancelled("relay setup cancelled"), + ); + + assert!(!registry.abort_relay("sbx-test", "old", &close)); + assert!(!registry.abort_relay("other", "other-session", &close)); + assert!( + registry + .pending_relays + .lock() + .unwrap() + .contains_key(&channel_id) + ); + assert!(matches!( + relay_rx.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + )); + + // Pending relays survive session replacement and can be replayed to + // the current session, which owns their cancellation before claim. + assert!(registry.abort_relay("sbx-test", "current", &close)); + assert_eq!( + relay_rx.try_recv().unwrap().unwrap_err().code(), + tonic::Code::Cancelled + ); + } + // ---- claim_relay: expiry, drop, wiring ---- #[test] @@ -3167,7 +3307,7 @@ mod tests { #[test] fn claim_relay_receiver_dropped_returns_internal() { let registry = SupervisorSessionRegistry::new(); - let (relay_tx, relay_rx) = oneshot::channel::>(); + let (relay_tx, relay_rx) = oneshot::channel::>(); drop(relay_rx); // Gateway-side waiter has given up already. registry.pending_relays.lock().unwrap().insert( "ch-1".to_string(), @@ -3183,7 +3323,7 @@ mod tests { #[tokio::test] async fn claim_relay_connects_both_ends() { let registry = SupervisorSessionRegistry::new(); - let (relay_tx, relay_rx) = oneshot::channel::>(); + let (relay_tx, relay_rx) = oneshot::channel::>(); registry.pending_relays.lock().unwrap().insert( "ch-io".to_string(), pending_relay("sbx-test", relay_tx, Instant::now()), diff --git a/crates/openshell-server/tests/supervisor_relay_integration.rs b/crates/openshell-server/tests/supervisor_relay_integration.rs index 5856acbbda..bb137935eb 100644 --- a/crates/openshell-server/tests/supervisor_relay_integration.rs +++ b/crates/openshell-server/tests/supervisor_relay_integration.rs @@ -669,7 +669,10 @@ async fn run_echo_supervisor(channel: Channel, channel_id: String) { out_tx .send(RelayFrame { payload: Some(openshell_core::proto::relay_frame::Payload::Init( - RelayInit { channel_id }, + RelayInit { + channel_id, + ..Default::default() + }, )), }) .await @@ -729,6 +732,40 @@ async fn relay_round_trips_bytes() { assert_eq!(&buf, b"hello relay"); } +#[tokio::test] +async fn legacy_relay_drains_reply_after_response_eof() { + use openshell_core::stream_lifecycle::{Frame, Payload}; + tokio::time::timeout(Duration::from_secs(5), async { + let registry = Arc::new(SupervisorSessionRegistry::new()); + let channel = spawn_gateway(Arc::clone(®istry)).await; + let mut session_rx = register_session(®istry, "sbx"); + let (channel_id, relay_rx) = registry.open_relay("sbx", Duration::from_secs(2)).await.unwrap(); + session_rx.recv().await.unwrap(); + let (input, rx) = mpsc::channel(4); + input.send(RelayFrame { + payload: Some(openshell_core::proto::relay_frame::Payload::Init(RelayInit { + channel_id, + ..Default::default() + })), + }).await.unwrap(); + let mut client = OpenShellClient::new(channel); + let mut response = client.relay_stream(ReceiverStream::new(rx)).await.unwrap().into_inner(); + let mut pipe = relay_rx.await.unwrap().unwrap(); + pipe.write_all(b"request").await.unwrap(); + pipe.shutdown().await.unwrap(); + assert!(matches!(response.message().await.unwrap().unwrap().payload(), Payload::Data(data) if data == b"request")); + assert!(response.message().await.unwrap().is_none()); + // A legacy supervisor drains target output after response EOF, sending + // the reply on the still-open HTTP/2 request stream. + input.send(RelayFrame::data(b"delayed reply".to_vec())).await.unwrap(); + drop(input); + let mut reply = Vec::new(); + pipe.read_to_end(&mut reply).await.unwrap(); + assert_eq!(reply, b"delayed reply"); + pipe.completed().await.unwrap(); + }).await.expect("legacy relay drain hung"); +} + #[tokio::test] async fn relay_closes_cleanly_when_gateway_drops() { let registry = Arc::new(SupervisorSessionRegistry::new()); @@ -777,7 +814,10 @@ async fn relay_sees_eof_when_supervisor_closes() { out_tx .send(RelayFrame { payload: Some(openshell_core::proto::relay_frame::Payload::Init( - RelayInit { channel_id }, + RelayInit { + channel_id, + ..Default::default() + }, )), }) .await @@ -913,3 +953,43 @@ async fn test_health_store() -> Arc { .expect("connect in-memory sqlite store for tests"), ) } + +/// A negotiated FIN must not become terminal success before the opposite +/// direction finishes, including when a typed abort follows the FIN. +#[tokio::test] +async fn negotiated_relay_fin_preserves_input_and_terminal_status() { + use openshell_core::stream_lifecycle::{self, Frame}; + tokio::time::timeout(Duration::from_secs(5), async { + for abort_after_fin in [false, true] { + let registry = Arc::new(SupervisorSessionRegistry::new()); + let channel = spawn_gateway(Arc::clone(®istry)).await; + let mut session_rx = register_session(®istry, "sbx"); + let (channel_id, relay_rx) = registry.open_relay("sbx", Duration::from_secs(2)).await.unwrap(); + session_rx.recv().await.unwrap(); + let (input, rx) = mpsc::channel(4); + input.send(RelayFrame { payload: Some(openshell_core::proto::relay_frame::Payload::Init(RelayInit { + channel_id, + capabilities: stream_lifecycle::capabilities(), + session_id: "sess-1".into(), + })) }).await.unwrap(); + let mut client = OpenShellClient::new(channel); + let mut response = client.relay_stream(ReceiverStream::new(rx)).await.unwrap().into_inner(); + let mut pipe = relay_rx.await.unwrap().unwrap(); + pipe.write_all(b"before FIN").await.unwrap(); + pipe.shutdown().await.unwrap(); + assert!(matches!(response.message().await.unwrap().unwrap().payload(), stream_lifecycle::Payload::Data(data) if data == b"before FIN")); + assert!(matches!(response.message().await.unwrap().unwrap().payload(), stream_lifecycle::Payload::HalfClose)); + if abort_after_fin { + pipe.abort_handle().abort(Status::deadline_exceeded("after FIN")); + assert_eq!(response.message().await.unwrap_err().code(), tonic::Code::DeadlineExceeded); + } else { + input.send(RelayFrame::data(b"after FIN".to_vec())).await.unwrap(); + drop(input); + let mut data = Vec::new(); + pipe.read_to_end(&mut data).await.unwrap(); + assert_eq!(data, b"after FIN"); + assert!(response.message().await.unwrap().is_none()); + } + } + }).await.expect("relay FIN test hung"); +} diff --git a/crates/openshell-supervisor-process/src/supervisor_session.rs b/crates/openshell-supervisor-process/src/supervisor_session.rs index 5be01017eb..fcff3c3d9e 100644 --- a/crates/openshell-supervisor-process/src/supervisor_session.rs +++ b/crates/openshell-supervisor-process/src/supervisor_session.rs @@ -10,9 +10,11 @@ //! and bridges bytes. The supervisor is a dumb byte bridge after target //! selection — it has no protocol awareness of the bytes flowing through. +use openshell_core::stream_lifecycle::{self, AbortHandle}; +use std::collections::HashMap; use std::net::IpAddr; -use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex}; use std::time::Duration; use openshell_core::proto::open_shell_client::OpenShellClient; @@ -26,9 +28,8 @@ use openshell_ocsf::{ ActivityId, BaseEventBuilder, ConnectionInfo, Endpoint, EventContext, NetworkActivityBuilder, OcsfEvent, SeverityId, StatusId, ocsf_emit, }; -use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; +use tokio::io::{AsyncRead, AsyncWrite}; use tokio::sync::{mpsc, watch}; -use tokio_stream::StreamExt; use tracing::{debug, warn}; use openshell_core::grpc_client; @@ -230,11 +231,6 @@ fn relay_close_from_gateway_event(ctx: &EventContext, channel_id: &str, reason: .build() } -/// Size of chunks read from the local SSH socket when forwarding bytes back -/// to the gateway over the gRPC response stream. 16 KiB matches the default -/// HTTP/2 frame size so each `RelayFrame::data` fits in one frame. -const RELAY_CHUNK_SIZE: usize = 16 * 1024; - trait TargetStream: AsyncRead + AsyncWrite + Send + Unpin {} impl TargetStream for T where T: AsyncRead + AsyncWrite + Send + Unpin {} @@ -415,6 +411,7 @@ async fn run_single_session( instance_id: config.instance_id.clone(), connection_epoch, supports_provider_readiness: true, + capabilities: stream_lifecycle::capabilities(), })), }) .await @@ -441,6 +438,9 @@ async fn run_single_session( _ => return Err("expected SessionAccepted or SessionRejected".into()), }; + let half_close = stream_lifecycle::supports_half_close(&accepted.capabilities); + let relays = SessionRelays::default(); + let heartbeat_secs = accepted .heartbeat_interval .as_ref() @@ -481,6 +481,9 @@ async fn run_single_session( channel: &channel, tx: &tx, terminating: &config.terminating, + half_close, + session_id: &accepted.session_id, + relays: &relays.0, }; handle_gateway_message( &msg, @@ -528,12 +531,13 @@ pub(crate) async fn test_bridge_ssh_relay( inbound: mpsc::Receiver>, out_tx: mpsc::Sender, ) { - let _ = bridge_relay( - Box::new(target), + let (read, write) = tokio::io::split(target); + let _ = stream_lifecycle::client_relay( tokio_stream::wrappers::ReceiverStream::new(inbound), + read, + write, out_tx, - "half-open-test".into(), - Arc::new(AtomicBool::new(false)), + false, ) .await; } @@ -557,7 +561,16 @@ pub async fn finalize_main_process_exit( Ok(()) } +#[derive(Default)] +struct SessionRelays(Arc>>); +// The map belongs to one control session, but established data RPCs retain it +// until completion. Losing the control stream alone must not abort those RPCs; +// a new session gets its own map and cannot cancel an older session's channels. + struct GatewayMessageContext<'a> { + half_close: bool, + session_id: &'a str, + relays: &'a Arc>>, sandbox_id: &'a str, ssh_socket_path: &'a std::path::Path, port_forward: &'a Arc, @@ -582,23 +595,43 @@ fn handle_gateway_message(msg: &GatewayMessage, context: &GatewayMessageContext< let port_forward = context.port_forward.clone(); let expected_ssh_peer_pid = context.expected_ssh_peer_pid; let terminating = Arc::clone(context.terminating); + let half_close = context.half_close; + let session_id = context.session_id.to_string(); + let relays = Arc::clone(context.relays); + let abort = AbortHandle::default(); + { + let mut active = relays.lock().unwrap(); + if active.contains_key(&channel_id) { + warn!(%channel_id, "duplicate relay open"); + return; + } + active.insert(channel_id.clone(), abort.clone()); + } let event = relay_open_event(openshell_ocsf::ctx::ctx(), &relay_open, &ssh_socket_path); ocsf_emit!(event); tokio::spawn(async move { let event_open = relay_open.clone(); - match handle_relay_open( - relay_open, - &ssh_socket_path, - port_forward, + let context = GatewayMessageContext { + sandbox_id: &sandbox_id, + ssh_socket_path: &ssh_socket_path, + port_forward: &port_forward, expected_ssh_peer_pid, - channel, - tx, - terminating, - ) - .await - { + channel: &channel, + tx: &tx, + terminating: &terminating, + half_close, + session_id: &session_id, + relays: &relays, + }; + let result: Result<(), Box> = tokio::select! { + biased; + error = abort.aborted() => Err(error.into()), + result = handle_relay_open(relay_open, &context) => result, + }; + relays.lock().unwrap().remove(&channel_id); + match result { Ok(()) => { let event = relay_closed_event( openshell_ocsf::ctx::ctx(), @@ -608,6 +641,15 @@ fn handle_gateway_message(msg: &GatewayMessage, context: &GatewayMessageContext< ocsf_emit!(event); } Err(e) => { + let status = e + .downcast_ref::() + .cloned() + .unwrap_or_else(|| tonic::Status::unavailable(e.to_string())); + let _ = tx.try_send(SupervisorMessage { + payload: Some(supervisor_message::Payload::RelayClose( + stream_lifecycle::close_message(channel_id.clone(), &status), + )), + }); let event = relay_failed_event( openshell_ocsf::ctx::ctx(), &event_open, @@ -626,6 +668,9 @@ fn handle_gateway_message(msg: &GatewayMessage, context: &GatewayMessageContext< }); } Some(gateway_message::Payload::RelayClose(close)) => { + if let Some(abort) = context.relays.lock().unwrap().get(&close.channel_id) { + abort.abort(stream_lifecycle::close_status(close)); + } let event = relay_close_from_gateway_event( openshell_ocsf::ctx::ctx(), &close.channel_id, @@ -647,32 +692,27 @@ fn handle_gateway_message(msg: &GatewayMessage, context: &GatewayMessageContext< /// frames carry raw SSH bytes in `data`. async fn handle_relay_open( relay_open: RelayOpen, - ssh_socket_path: &std::path::Path, - port_forward: Arc, - expected_ssh_peer_pid: Option, - channel: grpc_client::AuthedChannel, - tx: mpsc::Sender, - terminating: Arc, + context: &GatewayMessageContext<'_>, ) -> Result<(), Box> { let channel_id = relay_open.channel_id.clone(); let target = match open_target( &relay_open, - ssh_socket_path, - &port_forward, - expected_ssh_peer_pid, + context.ssh_socket_path, + context.port_forward, + context.expected_ssh_peer_pid, ) .await { Ok(target) => target, Err(err) => { - send_relay_open_result(&tx, &channel_id, false, err.to_string()).await; + send_relay_open_result(context.tx, &channel_id, false, err.to_string()).await; return Err(err); } }; - send_relay_open_result(&tx, &channel_id, true, String::new()).await; + send_relay_open_result(context.tx, &channel_id, true, String::new()).await; - let mut client = OpenShellClient::new(channel); + let mut client = OpenShellClient::new(context.channel.clone()); // Outbound chunks to the gateway. let (out_tx, out_rx) = mpsc::channel::(16); @@ -684,6 +724,12 @@ async fn handle_relay_open( payload: Some(openshell_core::proto::relay_frame::Payload::Init( RelayInit { channel_id: channel_id.clone(), + session_id: context.session_id.to_string(), + capabilities: if context.half_close { + stream_lifecycle::capabilities() + } else { + Vec::new() + }, }, )), }) @@ -693,7 +739,7 @@ async fn handle_relay_open( // Initiate the RPC. This rides the existing HTTP/2 connection. let response = match client.relay_stream(outbound).await { Ok(response) => response, - Err(e) if expected_transport_close_during_shutdown(&e, &terminating) => { + Err(e) if expected_transport_close_during_shutdown(&e, context.terminating) => { debug!( channel_id = %channel_id, error = %e, @@ -701,100 +747,18 @@ async fn handle_relay_open( ); return Ok(()); } - Err(e) => return Err(format!("relay_stream RPC failed: {e}").into()), + Err(e) => return Err(e.into()), }; - bridge_relay( - target, + let (read, write) = tokio::io::split(target); + stream_lifecycle::client_relay( response.into_inner(), + read, + write, out_tx, - channel_id, - terminating, + context.half_close, ) .await -} - -/// Forward the relay's data frames without interpreting the target protocol. -async fn bridge_relay( - target: Box, - mut inbound: impl tokio_stream::Stream> + Unpin, - out_tx: mpsc::Sender, - channel_id: String, - terminating: Arc, -) -> Result<(), Box> { - // Connect to the local SSH daemon on its Unix socket. - let (mut target_r, mut target_w) = tokio::io::split(target); - - debug!( - channel_id = %channel_id, - "relay bridge: connected to local target" - ); - - // Target → gRPC (out_tx): read local target, forward as `RelayFrame::data`. - let out_tx_writer = out_tx.clone(); - let target_to_grpc = tokio::spawn(async move { - let mut buf = vec![0u8; RELAY_CHUNK_SIZE]; - loop { - match target_r.read(&mut buf).await { - Ok(0) | Err(_) => break, - Ok(n) => { - let chunk = RelayFrame { - payload: Some(openshell_core::proto::relay_frame::Payload::Data( - buf[..n].to_vec(), - )), - }; - if out_tx_writer.send(chunk).await.is_err() { - break; - } - } - } - } - }); - - // gRPC (inbound) → target: drain inbound chunks into the local target socket. - let mut inbound_err: Option = None; - while let Some(next) = inbound.next().await { - match next { - Ok(frame) => { - let Some(openshell_core::proto::relay_frame::Payload::Data(data)) = frame.payload - else { - inbound_err = Some("relay inbound received non-data frame".to_string()); - break; - }; - if data.is_empty() { - continue; - } - if let Err(e) = target_w.write_all(&data).await { - inbound_err = Some(format!("write to target failed: {e}")); - break; - } - } - Err(e) => { - if expected_transport_close_during_shutdown(&e, &terminating) { - debug!( - channel_id = %channel_id, - error = %e, - "relay bridge: inbound closed during local shutdown" - ); - } else { - inbound_err = Some(format!("relay inbound errored: {e}")); - } - break; - } - } - } - - // Half-close the target socket's write side so the service sees EOF. - let _ = target_w.shutdown().await; - - // Dropping out_tx closes the outbound gRPC stream, letting the gateway - // observe EOF on its side too. - drop(out_tx); - let _ = target_to_grpc.await; - - if let Some(e) = inbound_err { - return Err(e.into()); - } - Ok(()) + .map_err(Into::into) } async fn send_relay_open_result( @@ -945,6 +909,34 @@ mod target_tests { mod ocsf_event_tests { use super::*; + #[tokio::test] + async fn control_session_loss_preserves_existing_data_relay() { + let session = SessionRelays::default(); + let abort = AbortHandle::default(); + session + .0 + .lock() + .unwrap() + .insert("channel".into(), abort.clone()); + let relay_map = Arc::clone(&session.0); + drop(session); + let replacement = SessionRelays::default(); + assert!(!replacement.0.lock().unwrap().contains_key("channel")); + assert!( + tokio::time::timeout(Duration::from_millis(20), abort.aborted()) + .await + .is_err() + ); + // Data-plane cancellation still works through the original owner. + relay_map + .lock() + .unwrap() + .get("channel") + .unwrap() + .abort(tonic::Status::cancelled("data stream cancelled")); + assert_eq!(abort.aborted().await.code(), tonic::Code::Cancelled); + } + #[cfg(target_os = "linux")] struct UnusedLoopbackConnector; diff --git a/docs/index.yml b/docs/index.yml index 596ac87172..c3e7ff7948 100644 --- a/docs/index.yml +++ b/docs/index.yml @@ -129,6 +129,8 @@ navigation: path: sdk/api-errors.mdx - page: "Protobuf Time Types" path: sdk/protobuf-time-types.mdx + - page: "Forwarding Stream Lifecycle" + path: reference/stream-lifecycle.mdx - folder: security title: "Security" - section: "Upgrade Guides" diff --git a/docs/reference/stream-lifecycle.mdx b/docs/reference/stream-lifecycle.mdx new file mode 100644 index 0000000000..8c6cb48b37 --- /dev/null +++ b/docs/reference/stream-lifecycle.mdx @@ -0,0 +1,49 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "Forwarding Stream Lifecycle" +slug: reference/stream-lifecycle +description: "Directional closure, cancellation, and compatibility for TCP forwarding and supervisor relays." +position: 12 +--- + +`ForwardTcp`, `RelayStream`, and `PeerRelay` carry independent request and response byte streams. A directional close ends bytes in one direction. Cancellation aborts the whole operation. + +## Close input and receive output + +Send the initialization frame once, followed by data frames. Close the request stream to signal input EOF. Continue reading responses until the RPC finishes, including its final status. Closing input does not discard a delayed response. + +After initialization, an additional init, an empty or unknown payload, or a request `half_close` frame returns `INVALID_ARGUMENT`. An empty data frame remains valid. The server does not impose a new application-idle or half-closed timeout; use a deadline or cancel the RPC when you need a bound. + +## Negotiate response closure + +Advertise `stream-half-close-v1` in the init frame's `capabilities` to accept response `half_close` frames. Unknown capability strings are ignored. Supervisors advertise support in `SupervisorHello`, use the capabilities returned in `SessionAccepted`, and include that session ID with the negotiated `RelayInit`. Gateways negotiate each `PeerRelay` hop independently. + +A response `half_close` follows all response bytes and closes only that direction. You can continue sending input. Keep reading the RPC after that frame: success requires both directions to finish, and a later error status still fails the operation. Repeated response FIN or data after FIN is a protocol error. + +Gateways also require downstream support before emitting response FIN. A `PeerRelay` response confirms effective support with `openshell-stream-half-close: v1` in its initial metadata. Missing metadata or a legacy supervisor downgrades the upstream forwarding path to response EOF. This uses the existing RPC initialization, without an additional application-data handshake. + +`openshell forward service` translates this response frame into a write shutdown on the local TCP socket. The SSH stdio proxy retains legacy response EOF because its stdout is a process descriptor, not an independently closable socket direction. Existing SDK `Close` methods retain their meanings; this extension does not change curated SDK interfaces. + +## Abort a relay + +`RelayClose` on `ConnectSupervisor` aborts a channel owned by that sandbox and supervisor session. Its typed code distinguishes `CANCELLED`, `DEADLINE_EXCEEDED`, `INVALID_ARGUMENT`, and `UNAVAILABLE`. Unknown or unspecified codes map to `UNAVAILABLE`; the human-readable reason is diagnostic. + +If the current supervisor session closes a channel before its relay stream is established, the gateway immediately reports the close status to the waiting caller and frees the pending relay slot. + +An abort interrupts blocked I/O and can discard buffered bytes. The first observed abort wins. The data stream's final status remains authoritative; a control message that arrives after completion cannot revise it. If a transport disconnect prevents delivery of the typed reason, the transport error remains the fallback. + +Losing the control stream alone does not cancel established data relays. They continue until their data stream finishes or fails. A replacement control session cannot cancel channels owned by the previous session. This does not guarantee survival of a shared transport failure or gateway restart. + +## Compatibility and behavior changes + +The protobuf additions preserve existing field numbers. Servers send response FIN frames only to peers advertising support. Legacy internal relays keep draining target replies after input EOF. Forward-facing RPCs still use legacy response EOF to finish the operation, so a mixed-version path cannot preserve further requests after response EOF. Upgrade the CLI, gateways, and supervisor runtime to use independent directional closure throughout a routed connection. + +The following behavior changes need release review even though the wire additions are compatible: + +- Malformed post-init frames now fail explicitly instead of being ignored or treated as closure. +- Relay I/O and trailer errors now reach callers instead of appearing as successful EOF. +- `RelayClose` now cancels the owned operation instead of only recording an event. +- Negotiated streams can remain open after one direction closes, until the other finishes or the caller cancels. + +There is no automatic replay of application bytes or new reconnect guarantee. diff --git a/e2e/python/test_forward_lifecycle.py b/e2e/python/test_forward_lifecycle.py new file mode 100644 index 0000000000..da517f1994 --- /dev/null +++ b/e2e/python/test_forward_lifecycle.py @@ -0,0 +1,99 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import threading +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import TYPE_CHECKING + +import pytest + +from openshell._proto import openshell_pb2 + +if TYPE_CHECKING: + from collections.abc import Callable + + from openshell import Sandbox, SandboxClient + + +@pytest.mark.parametrize("response_first", [False, True]) +def test_forward_preserves_bytes_after_opposite_fin( + sandbox: Callable[..., Sandbox], + sandbox_client: SandboxClient, + response_first: bool, +) -> None: + """Exercise real TCP FIN through ForwardTcp and the built supervisor.""" + port_file = f"/sandbox/forward-{uuid.uuid4().hex}.port" + server = f""" +import socket +from pathlib import Path +with socket.socket() as listener: + listener.bind(('127.0.0.1', 0)) + listener.listen(1) + listener.settimeout(30) + Path({port_file!r}).write_text(str(listener.getsockname()[1])) + with listener.accept()[0] as connection: + connection.settimeout(30) + if {response_first!r}: + connection.sendall(b'response') + connection.shutdown(socket.SHUT_WR) + request = bytearray() + while chunk := connection.recv(4096): + request.extend(chunk) + assert request == b'request after FIN' + if not {response_first!r}: + connection.sendall(b'response') + connection.shutdown(socket.SHUT_WR) +""" + wait_ready = f""" +import time +from pathlib import Path +path = Path({port_file!r}) +deadline = time.monotonic() + 20 +while not path.exists() and time.monotonic() < deadline: + time.sleep(0.05) +print(path.read_text()) +""" + with sandbox(delete_on_exit=True) as sb, ThreadPoolExecutor(max_workers=1) as pool: + target = pool.submit(sb.exec, ["python", "-c", server], timeout_seconds=45) + ready = sb.exec(["python", "-c", wait_ready], timeout_seconds=25) + assert ready.exit_code == 0, ready.stderr + port = int(ready.stdout.strip()) + response_fin = threading.Event() + session_request = openshell_pb2.CreateSshSessionRequest(sandbox=sb.sandbox.name) + session_request.workspace_scope.workspace = "default" + session = sandbox_client._stub.CreateSshSession(session_request, timeout=15) + + def requests(): + yield openshell_pb2.TcpForwardFrame( + init=openshell_pb2.TcpForwardInit( + sandbox=sb.sandbox.name, + workspace="default", + authorization_token=session.token, + tcp=openshell_pb2.TcpRelayTarget(host="127.0.0.1", port=port), + capabilities=["stream-half-close-v1"], + ) + ) + if response_first: + assert response_fin.wait(30), "response FIN did not arrive" + yield openshell_pb2.TcpForwardFrame(data=b"request after FIN") + + output = bytearray() + fin_count = 0 + try: + for frame in sandbox_client._stub.ForwardTcp(requests(), timeout=35): + kind = frame.WhichOneof("payload") + if kind == "half_close": + fin_count += 1 + response_fin.set() + else: + assert kind == "data" and fin_count == 0 + output.extend(frame.data) + finally: + response_fin.set() + assert output == b"response" + assert fin_count == 1 + result = target.result(timeout=10) + assert result.exit_code == 0, result.stderr diff --git a/proto/openshell.proto b/proto/openshell.proto index 83ffbe75ad..d56699e96f 100644 --- a/proto/openshell.proto +++ b/proto/openshell.proto @@ -1940,6 +1940,8 @@ message ExecSandboxEvent { // Initial frame for one TCP forward stream. message TcpForwardInit { + // Advertised extensions; unknown capabilities are ignored. + repeated string capabilities = 8; string sandbox = 1; string workspace = 3; // Optional service identifier for audit/correlation. @@ -1954,11 +1956,19 @@ message TcpForwardInit { string authorization_token = 7 [(openshell.options.v1.secret) = true]; } +// Response-only FIN. Request EOF remains the input FIN. This is not cancellation. +message StreamHalfClose {} + // A single frame on the CLI-to-gateway TCP forward stream. +// Init must be first; subsequent request frames must contain data. Request EOF +// closes input while output drains. A response half_close is sent only when +// stream-half-close-v1 was advertised; final success waits for both directions. message TcpForwardFrame { oneof payload { TcpForwardInit init = 1; bytes data = 2; + // Response only, negotiated with stream-half-close-v1. No data may follow. + StreamHalfClose half_close = 3; } } @@ -3001,10 +3011,14 @@ message SupervisorHello { uint64 connection_epoch = 3; // The supervisor can report credential, policy, and launch-environment installation. bool supports_provider_readiness = 4; + // Supported relay extensions. + repeated string capabilities = 5; } // Gateway accepts the supervisor session. message SessionAccepted { + // Intersection of gateway and SupervisorHello capabilities. + repeated string capabilities = 3; reserved 2; reserved "heartbeat_interval_secs"; // Gateway-assigned session ID for this connection. @@ -3077,23 +3091,32 @@ message TcpRelayTarget { // Initial RelayStream frame sent by the supervisor to claim a pending relay. message RelayInit { + // Echo negotiated capabilities for this relay stream. + repeated string capabilities = 2; // Gateway-allocated channel identifier (UUID). string channel_id = 1; + // Owning ConnectSupervisor session; required with stream-half-close-v1. + string session_id = 3; } // A single frame on the RelayStream RPC. // // The supervisor MUST send `init` as the first frame. All subsequent frames -// in either direction carry raw bytes in `data`. +// on the request stream carry raw bytes in `data`; request EOF is FIN. +// Negotiated response half_close preserves the still-open request direction. message RelayFrame { oneof payload { RelayInit init = 1; bytes data = 2; + // Response only, negotiated with stream-half-close-v1. No data may follow. + StreamHalfClose half_close = 3; } } // Initial frame for gateway peer relay forwarding. message PeerRelayInit { + // Supported extensions for this gateway-to-gateway stream. + repeated string capabilities = 4; // Stable sandbox UUID whose supervisor relay should be opened. string sandbox_id = 1; // Relay target to ask the owning gateway to open on its local supervisor @@ -3108,6 +3131,8 @@ message PeerRelayFrame { oneof payload { PeerRelayInit init = 1; bytes data = 2; + // Response only, negotiated with stream-half-close-v1. No data may follow. + StreamHalfClose half_close = 3; } } @@ -3121,8 +3146,20 @@ message RelayOpenResult { string error = 3; } -// Either side requests closure of a relay channel. +// Abort reason. Unknown/unspecified values map to UNAVAILABLE. +enum RelayCloseCode { + RELAY_CLOSE_CODE_UNSPECIFIED = 0; + RELAY_CLOSE_CODE_CANCELLED = 1; + RELAY_CLOSE_CODE_DEADLINE_EXCEEDED = 2; + RELAY_CLOSE_CODE_INVALID_ARGUMENT = 3; + RELAY_CLOSE_CODE_UNAVAILABLE = 4; +} + +// Either side aborts a relay; buffered data is not guaranteed to drain. +// Graceful directional closure uses data-stream EOF/half_close instead. message RelayClose { + // Typed terminal outcome; reason is diagnostic only. + RelayCloseCode code = 3; // Channel identifier to close. string channel_id = 1; // Optional reason for closure. diff --git a/skills/openshell-cli/SKILL.md b/skills/openshell-cli/SKILL.md index 3977018a01..3773d6f5ea 100644 --- a/skills/openshell-cli/SKILL.md +++ b/skills/openshell-cli/SKILL.md @@ -872,6 +872,12 @@ runtime validation failure. ## Workflow 10: Service Access +For protocols that send EOF before receiving a reply, use `forward service` +and keep reading after closing the local socket's write direction. Independent +response FIN requires support throughout the CLI/gateway/supervisor path. See +[forwarding stream lifecycle](https://docs.nvidia.com/openshell/latest/reference/stream-lifecycle.md) +for mixed-version fallback, deadlines, and cancellation behavior. + Use `forward` for local access and `service` for a gateway-managed HTTP endpoint: ```bash