diff --git a/apps/aether-gateway/src/tunnel/embedded/body.rs b/apps/aether-gateway/src/tunnel/embedded/body.rs new file mode 100644 index 000000000..db196b8f8 --- /dev/null +++ b/apps/aether-gateway/src/tunnel/embedded/body.rs @@ -0,0 +1,186 @@ +use std::collections::VecDeque; +use std::sync::Arc; + +use bytes::{Bytes, BytesMut}; +use parking_lot::Mutex; +use tokio::sync::Notify; + +const CHUNK_BYTES: usize = 32 * 1024; + +#[derive(Debug)] +pub enum LocalBodyEvent { + Chunk(Bytes), + End, + Error(String), +} + +#[derive(Default)] +struct BufferState { + chunks: VecDeque, + bytes: usize, + terminal: Option>, + receiver_taken: bool, + receiver_closed: bool, +} + +pub(super) struct ResponseBuffer { + state: Mutex, + notify: Notify, + capacity: usize, +} + +impl ResponseBuffer { + pub(super) fn new(capacity: usize) -> Arc { + Arc::new(Self { + state: Mutex::new(BufferState::default()), + notify: Notify::new(), + capacity, + }) + } + + pub(super) fn take_receiver(self: &Arc) -> Option { + let mut state = self.state.lock(); + if state.receiver_taken { + return None; + } + state.receiver_taken = true; + Some(BodyReceiver { + buffer: Arc::clone(self), + finished: false, + }) + } + + pub(super) fn push(&self, mut payload: Bytes) -> bool { + let mut state = self.state.lock(); + if state.terminal.is_some() + || state.receiver_closed + || payload.len() > self.capacity.saturating_sub(state.bytes) + { + return false; + } + state.bytes += payload.len(); + while !payload.is_empty() { + if let Some(tail) = state + .chunks + .back_mut() + .filter(|chunk| chunk.len() < CHUNK_BYTES) + { + let count = payload.len().min(CHUNK_BYTES - tail.len()); + tail.extend_from_slice(&payload.split_to(count)); + } else { + let count = payload.len().min(CHUNK_BYTES); + let chunk = payload.split_to(count); + state.chunks.push_back( + chunk + .try_into_mut() + .unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())), + ); + } + } + drop(state); + self.notify.notify_waiters(); + true + } + + pub(super) fn finish(&self, result: Result<(), String>) { + let mut state = self.state.lock(); + if state.terminal.is_none() { + state.terminal = Some(result); + } + drop(state); + self.notify.notify_waiters(); + } +} + +pub(super) struct BodyReceiver { + buffer: Arc, + finished: bool, +} + +impl BodyReceiver { + pub(super) async fn recv(&mut self) -> Option { + if self.finished { + return None; + } + loop { + let notified = self.buffer.notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + { + let mut state = self.buffer.state.lock(); + if let Some(chunk) = state.chunks.pop_front() { + state.bytes -= chunk.len(); + return Some(LocalBodyEvent::Chunk(chunk.freeze())); + } + if let Some(terminal) = state.terminal.take() { + self.finished = true; + state.receiver_closed = true; + return Some(match terminal { + Ok(()) => LocalBodyEvent::End, + Err(error) => LocalBodyEvent::Error(error), + }); + } + } + notified.await; + } + } +} + +impl Drop for BodyReceiver { + fn drop(&mut self) { + let mut state = self.buffer.state.lock(); + state.receiver_closed = true; + state.chunks.clear(); + state.bytes = 0; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn error_survives_a_full_buffer() { + let buffer = ResponseBuffer::new(CHUNK_BYTES); + let mut receiver = buffer.take_receiver().unwrap(); + assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES]))); + buffer.finish(Err("proxy disconnected".into())); + assert!(matches!( + receiver.recv().await, + Some(LocalBodyEvent::Chunk(_)) + )); + assert!( + matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected") + ); + assert!(receiver.recv().await.is_none()); + } + + #[tokio::test] + async fn small_frames_are_coalesced_within_the_byte_budget() { + let buffer = ResponseBuffer::new(4096); + let mut receiver = buffer.take_receiver().unwrap(); + for _ in 0..4096 { + assert!(buffer.push(Bytes::from_static(b"x"))); + } + assert!(!buffer.push(Bytes::from_static(b"x"))); + assert_eq!(buffer.state.lock().chunks.len(), 1); + buffer.finish(Ok(())); + assert!( + matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096) + ); + assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End))); + } + + #[tokio::test] + async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() { + let buffer = ResponseBuffer::new(1024); + let mut receiver = buffer.take_receiver().unwrap(); + let task = tokio::spawn(async move { receiver.recv().await }); + tokio::task::yield_now().await; + buffer.finish(Err("cancelled".into())); + buffer.finish(Ok(())); + assert!( + matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled") + ); + } +} diff --git a/apps/aether-gateway/src/tunnel/embedded/flow_control_tests.rs b/apps/aether-gateway/src/tunnel/embedded/flow_control_tests.rs new file mode 100644 index 000000000..5e5165ef8 --- /dev/null +++ b/apps/aether-gateway/src/tunnel/embedded/flow_control_tests.rs @@ -0,0 +1,250 @@ +use super::*; + +async fn fixture( + window: u32, + capacity: usize, +) -> ( + Arc, + Arc, + Arc, + aether_runtime::BoundedQueueReceiver, +) { + let hub = HubRouter::new(ControlPlaneClient::disabled()); + let (sender, receiver) = bounded_queue(capacity); + let (close_tx, _) = watch::channel(false); + let connection = Arc::new( + ProxyConn::new( + 99, + "flow-test".into(), + "flow-test".into(), + sender, + close_tx, + 16, + 3, + ) + .with_settings(protocol::SettingsPayload { + initial_stream_window_bytes: window, + min_window_update_bytes: (window / 4).max(1), + drain_deadline_ms: 1000, + }), + ); + hub.register_proxy(Arc::clone(&connection)); + let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap(); + (hub, connection, stream, receiver) +} + +fn meta() -> protocol::RequestMeta { + protocol::RequestMeta { + provider_id: None, + endpoint_id: None, + key_id: None, + method: "GET".into(), + url: "https://example.com".into(), + headers: HashMap::new(), + stream: true, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 30, + follow_redirects: None, + http1_only: false, + transport_profile: None, + } +} + +async fn headers(hub: &Arc, stream: &LocalStream) { + let payload = serde_json::to_vec(&protocol::ResponseMeta { + status: 200, + headers: vec![], + }) + .unwrap(); + let mut frame = protocol::encode_frame( + stream.proxy_stream_id, + protocol::RESPONSE_HEADERS, + 0, + &payload, + ); + hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame) + .await; +} + +#[tokio::test] +async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() { + let (hub, _, stream, mut outbound) = fixture(128, 1).await; + headers(&hub, &stream).await; + assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64]))); + let mut receiver = stream.take_body_receiver().unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(10), receiver.recv()) + .await + .is_err() + ); + assert_eq!(*stream.response_consumed_since_update.lock(), 64); + outbound.recv().await.unwrap(); + let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv()) + .await + .unwrap(); + assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64)); + assert_eq!(*stream.response_consumed_since_update.lock(), 0); + let Message::Binary(data) = outbound.recv().await.unwrap() else { + panic!("expected binary update") + }; + let frame = aether_contracts::tunnel::Frame::decode(data).unwrap(); + let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap(); + assert_eq!( + frame.msg_type, + aether_contracts::tunnel::MsgType::WindowUpdate + ); + assert_eq!(update.delta_bytes, 64); + hub.cancel_local_stream(stream.id, "test complete"); +} + +#[tokio::test] +async fn response_credit_is_not_returned_until_consumed() { + let (hub, _, stream, mut outbound) = fixture(128, 4).await; + outbound.recv().await.unwrap(); + let mut body = protocol::encode_frame( + stream.proxy_stream_id, + protocol::RESPONSE_BODY, + 0, + &[b'x'; 128], + ); + hub.handle_proxy_frame(99, &mut body).await; + assert!(outbound.try_recv().is_err()); + let mut receiver = stream.take_body_receiver().unwrap(); + assert!( + matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128) + ); + assert!(outbound.try_recv().is_ok()); + hub.cancel_local_stream(stream.id, "test complete"); +} + +#[tokio::test] +async fn cancelled_stream_open_releases_slot_without_resetting_connection() { + let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await; + let opening_hub = Arc::clone(&hub); + let opening = + tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await }); + tokio::time::timeout(Duration::from_secs(1), async { + while hub.local_streams.len() != 2 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + opening.abort(); + assert!(matches!(opening.await, Err(error) if error.is_cancelled())); + assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1); + assert_eq!(hub.local_streams.len(), 1); + assert_eq!(hub.proxy_to_local.len(), 1); + assert!(connection.is_available()); + outbound.recv().await.unwrap(); + assert!(outbound.try_recv().is_err()); + let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap(); + outbound.recv().await.unwrap(); + hub.cancel_local_stream(first_stream.id, "test complete"); + outbound.recv().await.unwrap(); + hub.cancel_local_stream(next_stream.id, "test complete"); +} + +#[tokio::test] +async fn full_response_buffer_preserves_disconnect_error() { + let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await; + outbound.recv().await.unwrap(); + headers(&hub, &stream).await; + let mut receiver = stream.take_body_receiver().unwrap(); + for _ in 0..128 { + let mut frame = protocol::encode_frame( + stream.proxy_stream_id, + protocol::RESPONSE_BODY, + 0, + &vec![b'x'; 32 * 1024], + ); + hub.handle_proxy_frame(99, &mut frame).await; + } + hub.unregister_proxy(connection.id, &connection.node_id); + let mut bytes = 0; + loop { + match receiver.recv().await { + Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(), + Some(LocalBodyEvent::Error(error)) => { + assert!(error.contains("disconnected")); + break; + } + event => panic!("disconnect must not become normal EOF: {event:?}"), + } + } + assert_eq!(bytes, 4 * 1024 * 1024); +} + +#[tokio::test] +async fn slow_stream_does_not_block_another_stream_on_the_same_connection() { + let (hub, _, slow, mut outbound) = fixture(128, 512).await; + outbound.recv().await.unwrap(); + let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap(); + assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128]))); + let mut overflowing = protocol::encode_frame( + slow.proxy_stream_id, + protocol::RESPONSE_BODY, + 0, + b"overflow", + ); + tokio::time::timeout(Duration::from_secs(1), async { + hub.handle_proxy_frame(99, &mut overflowing).await; + headers(&hub, &fast).await; + assert_eq!( + fast.wait_headers(Duration::from_secs(1)) + .await + .unwrap() + .status, + 200 + ); + }) + .await + .expect("slow stream must not block connection reader"); + assert!(!hub.local_streams.contains_key(&slow.id)); + assert!(hub.local_streams.contains_key(&fast.id)); + hub.cancel_local_stream(fast.id, "test complete"); +} + +#[tokio::test] +async fn cancelling_a_stream_wakes_request_window_waiters() { + let (_, _, stream, _) = fixture(128, 512).await; + *stream.request_window.available.lock() = 0; + let waiter = tokio::spawn({ + let stream = Arc::clone(&stream); + async move { + stream + .acquire_request_window(1, Duration::from_secs(30)) + .await + } + }); + tokio::task::yield_now().await; + stream.fail("cancelled"); + assert!(tokio::time::timeout(Duration::from_secs(1), waiter) + .await + .unwrap() + .unwrap() + .is_err()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() { + for index in 0..256 { + let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1)); + let window = Arc::new(StreamFlowWindow::new(0)); + let waiter = tokio::spawn({ + let stream = Arc::clone(&stream); + let window = Arc::clone(&window); + async move { + stream.wait_headers(Duration::from_secs(1)).await.unwrap(); + window.acquire(1, Duration::from_secs(1)).await.unwrap(); + } + }); + stream.set_response_headers(protocol::ResponseMeta { + status: 200, + headers: vec![], + }); + window.add(1); + waiter.await.unwrap(); + } +} diff --git a/apps/aether-gateway/src/tunnel/embedded/hub.rs b/apps/aether-gateway/src/tunnel/embedded/hub.rs index 69101027c..7cb319bae 100644 --- a/apps/aether-gateway/src/tunnel/embedded/hub.rs +++ b/apps/aether-gateway/src/tunnel/embedded/hub.rs @@ -12,10 +12,11 @@ use axum::extract::ws::Message; use bytes::Bytes; use dashmap::DashMap; use parking_lot::{Mutex, RwLock}; -use tokio::sync::mpsc; use tokio::sync::{watch, Notify}; use tracing::{debug, info, warn}; +pub use super::body::LocalBodyEvent; +use super::body::{BodyReceiver, ResponseBuffer}; use super::control_plane::ControlPlaneClient; use super::protocol; @@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000; const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024; const CONNECTION_WARMUP: Duration = Duration::from_secs(1); +#[cfg(test)] +#[path = "flow_control_tests.rs"] +mod flow_control_tests; + static STREAM_INITIAL_WINDOW_BYTES: LazyLock = LazyLock::new(|| { std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES") .ok() @@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock = LazyLock::new(|| { .unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY) }); -static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock = LazyLock::new(|| { - STREAM_INITIAL_WINDOW_BYTES - .saturating_div(4) - .clamp(1, 1024 * 1024) -}); +pub(super) fn local_settings() -> protocol::SettingsPayload { + protocol::SettingsPayload { + initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES) + .min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32), + min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES + .saturating_div(4) + .clamp(1, 1024 * 1024), + drain_deadline_ms: *DRAIN_DEADLINE_MS, + } +} #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum SendStatus { @@ -91,6 +101,7 @@ impl ConnHealthState { struct StreamFlowWindow { available: Mutex, notify: Notify, + closed: AtomicBool, } impl StreamFlowWindow { @@ -98,6 +109,7 @@ impl StreamFlowWindow { Self { available: Mutex::new(u64::from(initial)), notify: Notify::new(), + closed: AtomicBool::new(false), } } @@ -109,6 +121,12 @@ impl StreamFlowWindow { let requested = bytes as u64; let started_at = Instant::now(); loop { + let notified = self.notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + if self.closed.load(Ordering::Acquire) { + return Err(()); + } { let mut available = self.available.lock(); if *available >= requested { @@ -120,10 +138,7 @@ impl StreamFlowWindow { let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else { return Err(()); }; - if tokio::time::timeout(remaining, self.notify.notified()) - .await - .is_err() - { + if tokio::time::timeout(remaining, notified).await.is_err() { return Err(()); } } @@ -138,6 +153,11 @@ impl StreamFlowWindow { drop(available); self.notify.notify_waiters(); } + + fn close(&self) { + self.closed.store(true, Ordering::Release); + self.notify.notify_waiters(); + } } #[derive(Debug, Clone, Copy)] @@ -209,6 +229,10 @@ impl BoundedOutbound { pub fn snapshot(&self) -> QueueSnapshot { self.tx.snapshot() } + + pub(super) fn subscribe_close(&self) -> watch::Receiver { + self.close_tx.subscribe() + } } pub struct ProxyConn { @@ -231,6 +255,7 @@ pub struct ProxyConn { flow_window_blocked_ms: AtomicU64, write_latency_last_us: AtomicU64, write_latency_ewma_us: AtomicU64, + settings: Mutex, } impl ProxyConn { @@ -244,6 +269,7 @@ impl ProxyConn { protocol_version: u8, ) -> Self { Self { + settings: Mutex::new(local_settings()), id, node_id, node_name, @@ -271,6 +297,11 @@ impl ProxyConn { self } + pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self { + *self.settings.get_mut() = settings; + self + } + pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self { self.node_generation = tunnel_generation; self @@ -565,13 +596,6 @@ pub struct LocalResponseHead { pub headers: Vec<(String, String)>, } -#[derive(Debug)] -pub enum LocalBodyEvent { - Chunk(Bytes), - End, - Error(String), -} - #[derive(Debug, Default)] struct LocalWaitState { response: Option, @@ -585,10 +609,11 @@ pub struct LocalStream { proxy_stream_id: u32, request_window: StreamFlowWindow, response_consumed_since_update: Mutex, + min_window_update_bytes: u32, + response_connection: Mutex>>, wait_state: Mutex, headers_notify: Notify, - body_tx: mpsc::Sender, - body_rx: Mutex>>, + body: Arc, terminal: AtomicBool, } @@ -600,7 +625,6 @@ impl LocalStream { proxy_stream_id: u32, initial_window_bytes: u32, ) -> Self { - let (body_tx, body_rx) = mpsc::channel(128); Self { id, tunnel_generation, @@ -608,10 +632,11 @@ impl LocalStream { proxy_stream_id, request_window: StreamFlowWindow::new(initial_window_bytes), response_consumed_since_update: Mutex::new(0), + min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024), + response_connection: Mutex::new(None), wait_state: Mutex::new(LocalWaitState::default()), headers_notify: Notify::new(), - body_tx, - body_rx: Mutex::new(Some(body_rx)), + body: ResponseBuffer::new(initial_window_bytes as usize), terminal: AtomicBool::new(false), } } @@ -632,26 +657,51 @@ impl LocalStream { self.request_window.add(delta); } - fn response_window_update_delta(&self, bytes: usize) -> Option { - if bytes == 0 { - return None; + async fn flush_response_credit(&self) -> Result<(), String> { + if self.terminal.load(Ordering::Acquire) { + return Ok(()); } - - let mut consumed = self.response_consumed_since_update.lock(); - *consumed = consumed.saturating_add(bytes as u64); - let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES); - if *consumed < threshold { - return None; + let connection = self + .response_connection + .lock() + .as_ref() + .and_then(std::sync::Weak::upgrade); + let Some(connection) = connection else { + return Ok(()); + }; + if connection.protocol_version() < 3 { + return Ok(()); } - - let delta = (*consumed).min(u64::from(u32::MAX)) as u32; - *consumed = consumed.saturating_sub(u64::from(delta)); - Some(delta) + let delta = { + let consumed = self.response_consumed_since_update.lock(); + if *consumed < u64::from(self.min_window_update_bytes) { + return Ok(()); + } + (*consumed).min(u64::from(u32::MAX)) as u32 + }; + let frame = protocol::encode_window_update(self.proxy_stream_id, delta); + if connection + .send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT) + .await + == SendStatus::Queued + { + let mut consumed = self.response_consumed_since_update.lock(); + *consumed = consumed.saturating_sub(u64::from(delta)); + return Ok(()); + } + if self.terminal.load(Ordering::Acquire) { + return Ok(()); + } + connection.request_close(); + Err("proxy flow-control update failed".to_string()) } pub async fn wait_headers(&self, timeout: Duration) -> Result { tokio::time::timeout(timeout, async { loop { + let notified = self.headers_notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); let outcome = { let state = self.wait_state.lock(); if let Some(response) = &state.response { @@ -662,15 +712,20 @@ impl LocalStream { if let Some(error) = outcome { return Err(error); } - self.headers_notify.notified().await; + notified.await; } }) .await .map_err(|_| "timed out waiting for response headers".to_string())? } - pub fn take_body_receiver(&self) -> Option> { - self.body_rx.lock().take() + pub fn take_body_receiver(self: &Arc) -> Option { + self.body.take_receiver().map(|receiver| LocalBodyReceiver { + receiver, + stream: Arc::clone(self), + failed: false, + pending: None, + }) } fn set_response_headers(&self, meta: protocol::ResponseMeta) { @@ -690,21 +745,11 @@ impl LocalStream { } } - async fn push_body_chunk(&self, payload: Bytes) -> bool { + fn push_body_chunk(&self, payload: Bytes) -> bool { if self.terminal.load(Ordering::Acquire) { return false; } - // Use a timeout to prevent a slow consumer from blocking the shared - // proxy-connection reader (head-of-line blocking across streams). - match tokio::time::timeout( - Duration::from_secs(5), - self.body_tx.send(LocalBodyEvent::Chunk(payload)), - ) - .await - { - Ok(Ok(())) => true, - _ => false, - } + self.body.push(payload) } fn finish(&self) { @@ -722,7 +767,8 @@ impl LocalStream { if notify { self.headers_notify.notify_waiters(); } - let _ = self.body_tx.try_send(LocalBodyEvent::End); + self.request_window.close(); + self.body.finish(Ok(())); } fn fail(&self, error: impl Into) { @@ -742,7 +788,38 @@ impl LocalStream { if notify { self.headers_notify.notify_waiters(); } - let _ = self.body_tx.try_send(LocalBodyEvent::Error(error)); + self.request_window.close(); + self.body.finish(Err(error)); + } +} + +pub struct LocalBodyReceiver { + receiver: BodyReceiver, + stream: Arc, + failed: bool, + pending: Option, +} + +impl LocalBodyReceiver { + pub async fn recv(&mut self) -> Option { + if self.failed { + return None; + } + if self.pending.is_none() { + let event = self.receiver.recv().await?; + if let LocalBodyEvent::Chunk(chunk) = &event { + let mut consumed = self.stream.response_consumed_since_update.lock(); + *consumed = consumed.saturating_add(chunk.len() as u64); + } + self.pending = Some(event); + } + if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) { + if let Err(error) = self.stream.flush_response_credit().await { + self.failed = true; + return Some(LocalBodyEvent::Error(error)); + } + } + self.pending.take() } } @@ -765,6 +842,21 @@ pub struct HubRouter { drain_reasons: Mutex>, } +struct PendingStreamGuard<'router> { + hub: &'router HubRouter, + connection: &'router ProxyConn, + stream_id: u64, + committed: bool, +} + +impl Drop for PendingStreamGuard<'_> { + fn drop(&mut self) { + if !self.committed && self.hub.cleanup_local_stream(self.stream_id) { + self.connection.release_stream(); + } + } +} + struct NodeStatusEvent { node_id: String, authenticated_key: Option, @@ -1166,17 +1258,27 @@ impl HubRouter { // Frames encoded successfully -- now register the stream. let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed); - let local_stream = Arc::new(LocalStream::new( + let settings = proxy_conn.settings.lock().clone(); + let mut local_stream = LocalStream::new( local_stream_id, proxy_conn.node_generation.clone(), proxy_conn.id, proxy_stream_id, - *STREAM_INITIAL_WINDOW_BYTES, - )); + settings.initial_stream_window_bytes, + ); + local_stream.min_window_update_bytes = settings.min_window_update_bytes; + *local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn)); + let local_stream = Arc::new(local_stream); self.local_streams .insert(local_stream_id, local_stream.clone()); self.proxy_to_local .insert((proxy_conn.id, proxy_stream_id), local_stream_id); + let mut pending_stream = PendingStreamGuard { + hub: self, + connection: &proxy_conn, + stream_id: local_stream_id, + committed: false, + }; let send_status = proxy_conn .send_wait( @@ -1195,10 +1297,11 @@ impl HubRouter { "open_local_stream dispatched" ); match send_status { - SendStatus::Queued => Ok(local_stream), + SendStatus::Queued => { + pending_stream.committed = true; + Ok(local_stream) + } SendStatus::Closed | SendStatus::Congested => { - self.cleanup_local_stream(local_stream_id); - proxy_conn.release_stream(); Err("proxy connection congested".to_string()) } } @@ -1243,7 +1346,9 @@ impl HubRouter { .map(|entry| entry.value().clone()) .ok_or_else(|| "proxy connection unavailable".to_string())?; - let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE); + let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE + .min(proxy_conn.settings.lock().initial_stream_window_bytes as usize); + let total_chunks = payload.len().div_ceil(chunk_size); let result = if total_chunks == 0 { if end_stream { self.send_request_body_frame(&proxy_conn, &stream, &[], true) @@ -1252,7 +1357,7 @@ impl HubRouter { Ok(()) } } else { - for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() { + for (index, chunk) in payload.chunks(chunk_size).enumerate() { let is_last_chunk = index + 1 == total_chunks; if let Err(error) = self .send_request_body_frame( @@ -1352,17 +1457,20 @@ impl HubRouter { } else { protocol::encode_stream_error(stream.proxy_stream_id, reason) }; - let _ = pc.send(Message::Binary(frame.into())); + if pc.send(Message::Binary(frame.into())) != SendStatus::Queued { + pc.request_close(); + } } stream.fail(reason.to_string()); } - fn cleanup_local_stream(&self, local_stream_id: u64) { + fn cleanup_local_stream(&self, local_stream_id: u64) -> bool { let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else { - return; + return false; }; self.proxy_to_local .remove(&(stream.proxy_conn_id, stream.proxy_stream_id)); + true } pub async fn handle_proxy_frame(self: &Arc, proxy_conn_id: u64, data: &mut [u8]) { @@ -1433,9 +1541,7 @@ impl HubRouter { .get(&proxy_conn_id) .map(|entry| entry.value().clone()); if let Some(pc) = pc { - let _ = pc - .send_wait(Message::Binary(pong.into()), Duration::from_millis(250)) - .await; + let _ = pc.send(Message::Binary(pong.into())); } } protocol::PONG => {} @@ -1509,11 +1615,36 @@ impl HubRouter { ); } protocol::SETTINGS => { - debug!( - msg_type = header.msg_type, - proxy_conn_id = proxy_conn_id, - "received tunnel protocol v3 SETTINGS from proxy" - ); + let settings = protocol::decode_payload_with_limit( + data, + &header, + MAX_TUNNEL_CONTROL_PAYLOAD_SIZE, + ) + .ok() + .and_then(|payload| { + serde_json::from_slice::(&payload).ok() + }) + .filter(|settings| settings.is_valid()); + if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) { + if header.stream_id != 0 || header.flags != 0 { + connection.request_close(); + return; + } + let Some(settings) = settings else { + connection.request_close(); + return; + }; + let local = local_settings(); + let settings = settings + .negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms); + let mut current = connection.settings.lock(); + if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings { + drop(current); + connection.request_close(); + return; + } + *current = settings; + } } protocol::WINDOW_UPDATE => { self.handle_window_update(proxy_conn_id, header.stream_id, data, &header); @@ -1739,19 +1870,8 @@ impl HubRouter { None => return, }; - let payload_len = payload.len(); - if !stream.push_body_chunk(Bytes::from(payload)).await { + if !stream.push_body_chunk(Bytes::from(payload)) { self.cancel_local_stream(local_id, "local relay response congested"); - return; - } - - if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) { - if pc.protocol_version() >= 3 { - if let Some(delta) = stream.response_window_update_delta(payload_len) { - let frame = protocol::encode_window_update(header.stream_id, delta); - let _ = pc.send(Message::Binary(frame.into())); - } - } } } diff --git a/apps/aether-gateway/src/tunnel/embedded/local_relay.rs b/apps/aether-gateway/src/tunnel/embedded/local_relay.rs index da6f23332..4dd6d49e7 100644 --- a/apps/aether-gateway/src/tunnel/embedded/local_relay.rs +++ b/apps/aether-gateway/src/tunnel/embedded/local_relay.rs @@ -9,14 +9,13 @@ use axum::body::{Body, Bytes}; use axum::extract::{ConnectInfo, Path, Request, State}; use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode}; use axum::response::IntoResponse; -use tokio::sync::mpsc; use tracing::warn; use crate::api::response::apply_streaming_response_headers; use crate::headers::should_skip_response_header; use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation; -use super::hub::{LocalBodyEvent, LocalStream}; +use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream}; use super::protocol; use super::{AppState, RelayRequestAuthenticated}; @@ -40,7 +39,7 @@ impl Drop for StreamGuard { pub(crate) struct DirectRelayResponse { status: u16, headers: Vec<(String, String)>, - body_rx: mpsc::Receiver, + body_rx: LocalBodyReceiver, request_guard: StreamGuard, _request_permit: Option, } @@ -55,10 +54,13 @@ impl DirectRelayResponse { } pub(crate) async fn next_chunk(&mut self) -> Result, String> { + if self.request_guard.finished { + return Ok(None); + } let event = self.body_rx.recv().await; match event { Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)), - Some(LocalBodyEvent::End) | None => { + Some(LocalBodyEvent::End) => { self.request_guard.finished = true; Ok(None) } @@ -66,6 +68,7 @@ impl DirectRelayResponse { self.request_guard.finished = true; Err(error) } + None => Err("tunnel response ended without a terminal frame".to_string()), } } } @@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream( .open_authorized_local_stream(node_id, &meta) .await .map_err(|error| format!("connect: {error}"))?; + let request_guard = StreamGuard { + hub: state.hub.clone(), + stream_id: stream.id, + finished: false, + }; if let Err(error) = state .hub .push_local_request_body(stream.id, body, true) @@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream( status: response_head.status, headers: response_head.headers, body_rx, - request_guard: StreamGuard { - hub: state.hub.clone(), - stream_id: stream.id, - finished: false, - }, + request_guard, _request_permit: request_permit, }) } @@ -259,6 +263,11 @@ pub async fn relay_request( ); } }; + let request_guard = StreamGuard { + hub: state.hub.clone(), + stream_id: stream.id, + finished: false, + }; let body_stream = match spool.body_stream().await { Ok(stream) => stream, Err(error) => { @@ -306,12 +315,6 @@ pub async fn relay_request( ); } - let request_guard = StreamGuard { - hub: state.hub.clone(), - stream_id: stream.id, - finished: false, - }; - let wait_timeout = relay_header_timeout(&meta); let response_head = match stream.wait_headers(wait_timeout).await { Ok(response) => response, @@ -373,6 +376,9 @@ pub async fn relay_request( } } } + if !guard.finished { + yield Err(io::Error::other("tunnel response ended without a terminal frame")); + } guard.finished = true; }; @@ -563,6 +569,100 @@ mod tests { request } + #[tokio::test] + async fn cancelled_relays_reset_streams_during_upload_and_header_wait() { + for direct in [true, false] { + for during_upload in [true, false] { + let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ + sample_connected_proxy_node("node-123"), + ])); + let data = Arc::new( + GatewayDataState::with_proxy_node_repository_for_tests(repository) + .with_system_config_values_for_tests( + Vec::<(String, serde_json::Value)>::new(), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let state = test_app_state().with_data(data); + let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8); + let (proxy_close_tx, _) = watch::channel(false); + let connection = Arc::new( + ProxyConn::new( + 500, + "node-123".into(), + "Node 123".into(), + proxy_tx, + proxy_close_tx, + 16, + 3, + ) + .with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string()) + .with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()) + .with_settings(protocol::SettingsPayload { + initial_stream_window_bytes: 128, + min_window_update_bytes: 32, + drain_deadline_ms: 1000, + }), + ); + state.hub.register_proxy(Arc::clone(&connection)); + let meta = protocol::RequestMeta { + provider_id: None, + endpoint_id: None, + key_id: None, + method: "POST".into(), + url: "https://example.com/".into(), + headers: HashMap::new(), + stream: true, + request_timeout_ms: None, + stream_first_byte_timeout_ms: None, + timeout: 30, + follow_redirects: None, + http1_only: false, + transport_profile: None, + }; + let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]); + let relay = tokio::spawn(async move { + if direct { + let _response = + super::open_direct_relay_stream(&state, "node-123", meta, body) + .await + .unwrap(); + } else { + let request = + authenticated_request(encode_relay_envelope(&meta, &body)).await; + let _response = relay_request( + Path("node-123".into()), + State(state), + ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))), + request, + ) + .await; + } + }); + recv_tunnel_test_frame(&mut proxy_rx, "request headers").await; + recv_tunnel_test_frame(&mut proxy_rx, "request body").await; + relay.abort(); + assert!(relay.await.unwrap_err().is_cancelled()); + let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await + else { + panic!("expected binary reset frame") + }; + let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap(); + assert_eq!( + frame.msg_type, + aether_contracts::tunnel::MsgType::ResetStream + ); + assert_eq!( + connection + .stream_count + .load(std::sync::atomic::Ordering::Relaxed), + 0 + ); + assert!(connection.is_available()); + } + } + } + #[test] fn relay_header_timeout_ignores_request_timeout_for_stream_requests() { let meta = protocol::RequestMeta { diff --git a/apps/aether-gateway/src/tunnel/embedded/mod.rs b/apps/aether-gateway/src/tunnel/embedded/mod.rs index ca570a1f0..ab69f2e9f 100644 --- a/apps/aether-gateway/src/tunnel/embedded/mod.rs +++ b/apps/aether-gateway/src/tunnel/embedded/mod.rs @@ -1,3 +1,4 @@ +mod body; mod control_plane; mod hub; mod local_relay; diff --git a/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs b/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs index 9834e510f..750d9583e 100644 --- a/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs +++ b/apps/aether-gateway/src/tunnel/embedded/proxy_conn.rs @@ -9,6 +9,7 @@ use aether_runtime::bounded_queue; use axum::extract::ws::{Message, WebSocket}; use futures_util::{SinkExt, StreamExt}; use tokio::sync::watch; +use tokio::task::JoinSet; use tracing::{debug, info, warn}; use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus}; @@ -84,6 +85,35 @@ pub async fn handle_proxy_connection( let (tx, mut rx) = bounded_queue::(cfg.outbound_queue_capacity); let (close_tx, mut close_rx) = watch::channel(false); + let settings = if protocol_version >= 3 { + let Some(settings) = read_proxy_settings( + &mut ws_tx, + &mut ws_rx, + security.as_deref(), + protocol_version, + ) + .await + else { + warn!(conn_id, "proxy SETTINGS negotiation failed"); + return; + }; + let local = super::hub::local_settings(); + let negotiated = + settings.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms); + let message = Message::Binary(protocol::encode_settings(&negotiated).into()); + let Ok(message) = encrypt_message(message, security.as_deref()) else { + return; + }; + if !matches!( + tokio::time::timeout(PROXY_HELLO_TIMEOUT, ws_tx.send(message)).await, + Ok(Ok(())) + ) { + return; + } + negotiated + } else { + super::hub::local_settings() + }; let conn = ProxyConn::new( conn_id, node_id.clone(), @@ -93,7 +123,8 @@ pub async fn handle_proxy_connection( max_streams, protocol_version, ) - .with_tunnel_generation(node_generation); + .with_tunnel_generation(node_generation) + .with_settings(settings); let conn = match (security_key.clone(), management_token_credential) { (Some(key), None) => Arc::new(conn.with_authenticated_key(key)), (None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)), @@ -416,19 +447,22 @@ async fn run_proxy_reader( let idle_enabled = !idle_timeout.is_zero(); let mut oversized_count = 0u32; let mut frames_received: u64 = 0; + let mut close_rx = conn.outbound.subscribe_close(); + let mut heartbeats = JoinSet::new(); loop { - let msg = if idle_enabled { - tokio::select! { - msg = ws_rx.next() => msg, - _ = tokio::time::sleep(idle_timeout) => { - warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout"); - let _ = conn.send(Message::Binary(protocol::encode_goaway().into())); - conn.request_close(); - break; - } + if conn.outbound.is_closing() { + break; + } + while heartbeats.try_join_next().is_some() {} + let msg = tokio::select! { + biased; + _ = close_rx.changed() => break, + msg = ws_rx.next() => msg, + _ = tokio::time::sleep(idle_timeout), if idle_enabled => { + warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout"); + conn.request_close(); + break; } - } else { - ws_rx.next().await }; match msg { @@ -463,7 +497,27 @@ async fn run_proxy_reader( continue; } - hub.handle_proxy_frame(conn.id, &mut data).await; + let is_heartbeat = protocol::FrameHeader::parse(&data) + .is_some_and(|header| header.msg_type == protocol::HEARTBEAT_DATA); + if is_heartbeat { + if heartbeats.is_empty() { + let heartbeat_hub = Arc::clone(&hub); + let conn_id = conn.id; + heartbeats.spawn(async move { + if tokio::time::timeout( + Duration::from_secs(10), + heartbeat_hub.handle_proxy_frame(conn_id, &mut data), + ) + .await + .is_err() + { + warn!(conn_id, "proxy heartbeat processing timed out"); + } + }); + } + } else { + hub.handle_proxy_frame(conn.id, &mut data).await; + } } Some(Ok(Message::Close(_))) | None => { info!( @@ -489,6 +543,56 @@ async fn run_proxy_reader( _ => {} } } + heartbeats.shutdown().await; +} + +async fn read_proxy_settings( + ws_tx: &mut futures_util::stream::SplitSink, + ws_rx: &mut futures_util::stream::SplitStream, + security: Option<&SecureFrameCodec>, + protocol_version: u8, +) -> Option { + tokio::time::timeout(PROXY_HELLO_TIMEOUT, async { + let mut hello_received = security.is_some(); + for _ in 0..MAX_PREAUTH_PINGS { + match ws_rx.next().await? { + Ok(Message::Binary(data)) => { + if data.len() > 256 * 1024 { + return None; + } + let data = decrypt_message(data, security).ok()?; + let frame = Frame::decode(data.into()).ok()?; + if frame.stream_id != 0 || frame.flags != 0 { + return None; + } + match frame.msg_type { + MsgType::Hello if !hello_received => { + let hello = + serde_json::from_slice::(&frame.payload).ok()?; + if hello.protocol_version != protocol_version { + return None; + } + hello_received = true; + } + MsgType::Settings if hello_received => { + let settings = + serde_json::from_slice::(&frame.payload) + .ok()?; + return settings.is_valid().then_some(settings); + } + _ => return None, + } + } + Ok(Message::Ping(payload)) => ws_tx.send(Message::Pong(payload)).await.ok()?, + Ok(Message::Pong(_)) => {} + _ => return None, + } + } + None + }) + .await + .ok() + .flatten() } fn encrypt_message( @@ -523,6 +627,158 @@ fn decrypt_message( #[cfg(test)] mod tests { + #[cfg(feature = "testkit")] + #[tokio::test] + async fn slow_heartbeat_does_not_block_response_frames() { + use std::sync::atomic::{AtomicUsize, Ordering}; + use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message as ClientMessage}; + + let called = Arc::new(AtomicUsize::new(0)); + let callback_called = Arc::clone(&called); + let control_plane = super::super::control_plane::ControlPlaneClient::local( + move |_, _| { + callback_called.fetch_add(1, Ordering::SeqCst); + Box::pin(std::future::pending()) + }, + |_, _, _, _| Box::pin(async { Ok(()) }), + ); + let data = crate::data::GatewayDataState::with_tunnel_management_auth_for_testkit( + "heartbeat-test", + "heartbeat-generation", + "ae-tunnel-harness-management-token", + aether_crypto::DEVELOPMENT_ENCRYPTION_KEY, + ) + .unwrap(); + let state = super::super::AppState::new( + control_plane, + ConnConfig { + ping_interval: Duration::from_secs(60), + idle_timeout: Duration::ZERO, + outbound_queue_capacity: 128, + }, + 16, + ) + .with_data(Arc::new(data)); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let router = super::super::build_router_with_state(state.clone()); + let server = tokio::spawn(async move { + axum::serve( + listener, + router.into_make_service_with_connect_info::(), + ) + .await + .unwrap(); + }); + let mut request = format!("ws://{address}/api/internal/proxy-tunnel") + .into_client_request() + .unwrap(); + let headers = request.headers_mut(); + headers.insert("x-node-id", "heartbeat-test".parse().unwrap()); + headers.insert( + aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER, + "heartbeat-generation".parse().unwrap(), + ); + headers.insert( + aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER, + "3".parse().unwrap(), + ); + headers.insert( + "authorization", + "Bearer ae-tunnel-harness-management-token".parse().unwrap(), + ); + let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap(); + let hello = HelloPayload { + protocol_version: 3, + capabilities: vec![], + session_id: None, + replica_id: None, + }; + websocket + .send(ClientMessage::Binary(protocol::encode_hello(&hello).into())) + .await + .unwrap(); + websocket + .send(ClientMessage::Binary( + protocol::encode_settings(&super::super::hub::local_settings()).into(), + )) + .await + .unwrap(); + let ClientMessage::Binary(settings) = websocket.next().await.unwrap().unwrap() else { + panic!("expected SETTINGS") + }; + assert_eq!(Frame::decode(settings).unwrap().msg_type, MsgType::Settings); + tokio::time::timeout(Duration::from_secs(1), async { + while !state.hub.has_local_proxy("heartbeat-test") { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + let meta: protocol::RequestMeta = serde_json::from_value(serde_json::json!({ + "method": "GET", "url": "https://example.com", "headers": {}, "stream": true, "timeout": 10 + })).unwrap(); + let stream = state + .hub + .open_local_stream("heartbeat-test", &meta) + .await + .unwrap(); + let ClientMessage::Binary(request) = websocket.next().await.unwrap().unwrap() else { + panic!("expected request headers") + }; + let stream_id = Frame::decode(request).unwrap().stream_id; + let heartbeat = Frame::control( + MsgType::HeartbeatData, + serde_json::to_vec(&serde_json::json!({"node_id": "heartbeat-test"})).unwrap(), + ); + websocket + .send(ClientMessage::Binary(heartbeat.encode())) + .await + .unwrap(); + tokio::time::timeout(Duration::from_secs(1), async { + while called.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + for _ in 0..8 { + websocket + .send(ClientMessage::Binary(heartbeat.encode())) + .await + .unwrap(); + } + let response = Frame::new( + stream_id, + MsgType::ResponseHeaders, + 0, + serde_json::to_vec(&serde_json::json!({"status": 200, "headers": []})).unwrap(), + ); + websocket + .send(ClientMessage::Binary(response.encode())) + .await + .unwrap(); + assert_eq!( + stream + .wait_headers(Duration::from_secs(1)) + .await + .unwrap() + .status, + 200 + ); + assert_eq!(called.load(Ordering::SeqCst), 1); + state.hub.request_close_all_proxies(); + tokio::time::timeout(Duration::from_secs(1), async { + while state.hub.has_local_proxy("heartbeat-test") { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + server.abort(); + let _ = server.await; + } + use super::*; const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="; diff --git a/apps/aether-tunnel/Cargo.toml b/apps/aether-tunnel/Cargo.toml index 387f7565b..e6f20bb76 100644 --- a/apps/aether-tunnel/Cargo.toml +++ b/apps/aether-tunnel/Cargo.toml @@ -47,3 +47,4 @@ uuid.workspace = true [dev-dependencies] aether-gateway = { workspace = true, features = ["testkit"] } +tokio = { version = "1", features = ["test-util"] } diff --git a/apps/aether-tunnel/README.md b/apps/aether-tunnel/README.md index 2d2895759..26b4c29ba 100644 --- a/apps/aether-tunnel/README.md +++ b/apps/aether-tunnel/README.md @@ -4,6 +4,15 @@ Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道 Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。 +## 流式传输与升级注意事项 + +- 协议 v3 连接在 `HELLO` / `SETTINGS` 协商后才接收业务请求。实际双向流窗口取 gateway 与 agent 配置的较小值,信用更新阈值不超过该窗口的四分之一;单帧也不会超过协商窗口。 +- 响应缓冲按字节限额并合并小帧,结束和错误状态独立保存。慢消费者不会阻塞同一隧道其他流的读取;超出窗口或缓冲预算的流会被明确终止,不会静默截断。 +- 信用更新在消费数据后可靠入队;启用重定向重放时,进入有界重放缓存也视为请求体消费。持续无法投递关键控制帧时会关闭连接并向在途请求报告错误。 +- 客户端取消会终止对应上游请求,断连会回收 session 的 writer、heartbeat 和请求任务。正常 drain 在配置期限内继续处理已有流,期限到达后终止残留任务。 +- 建议先升级 gateway,再升级 agent。既有 v3 agent 已发送 `HELLO` / `SETTINGS`,可连接新 gateway;自定义 v3 节点必须完成这两步握手。协议 v1/v2 保留旧握手。与旧 gateway 混用时应保持默认窗口配置,不能依赖旧 gateway 应用新的窗口协商。 +- 自动重连恢复后续请求,不会自动续传已经输出的 SSE,也不会无条件重放已经发送的请求。 + ## 安装 `aether-tunnel` 会根据宿主机自动选择服务管理器: diff --git a/apps/aether-tunnel/src/config.rs b/apps/aether-tunnel/src/config.rs index 31772fb4c..d7ee7f595 100644 --- a/apps/aether-tunnel/src/config.rs +++ b/apps/aether-tunnel/src/config.rs @@ -764,6 +764,13 @@ impl Config { if self.tunnel_stream_initial_window_bytes == 0 { anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0"); } + if u64::from(self.tunnel_stream_initial_window_bytes) + > aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64 + { + anyhow::bail!( + "tunnel_stream_initial_window_bytes exceeds the maximum tunnel payload size" + ); + } if self.tunnel_drain_deadline_ms == 0 { anyhow::bail!("tunnel_drain_deadline_ms must be > 0"); } diff --git a/apps/aether-tunnel/src/tunnel/client.rs b/apps/aether-tunnel/src/tunnel/client.rs index 5ccacc76e..49c593905 100644 --- a/apps/aether-tunnel/src/tunnel/client.rs +++ b/apps/aether-tunnel/src/tunnel/client.rs @@ -203,19 +203,34 @@ pub async fn connect_and_run( let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream); // Spawn writer task (with WebSocket ping keepalive) - let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security( + let (frame_tx, writer_handle) = writer::spawn_writer_with_metrics_and_security( ws_sink, ping_interval, Some(Arc::clone(&server.tunnel_metrics)), security.clone(), ); + let mut writer_handle = super::task::SessionTask::new(writer_handle); send_protocol_v3_hello(&frame_tx, &security_session, state).await; - let drain_signal = spawn_drain_signal( + let (session_drain_tx, session_drain_rx) = watch::channel(*drain.borrow()); + let forward_drain_tx = session_drain_tx.clone(); + let mut external_drain = drain; + let forward_drain = super::task::SessionTask::new(tokio::spawn(async move { + loop { + if *external_drain.borrow() { + let _ = forward_drain_tx.send(true); + break; + } + if external_drain.changed().await.is_err() { + break; + } + } + })); + let drain_signal = super::task::SessionTask::new(spawn_drain_signal( conn_idx, frame_tx.clone(), - drain.clone(), + session_drain_rx.clone(), state.config.tunnel_drain_deadline_ms, - ); + )); // Spawn heartbeat task (only for primary connection to avoid // resetting shared atomic metrics via swap(0)) @@ -237,16 +252,19 @@ pub async fn connect_and_run( // ensures we detect this and trigger a reconnect promptly. let state_clone = Arc::clone(state); let server_clone = Arc::clone(server); - let outcome = tokio::select! { - result = dispatcher::run_with_security( + let outcome = { + let dispatch = dispatcher::run_with_security( state_clone, server_clone, ws_read, frame_tx.clone(), hb_handle, - drain.clone(), + session_drain_rx, security.clone(), - ) => { + ); + tokio::pin!(dispatch); + tokio::select! { + result = &mut dispatch => { match result { Ok(()) => Ok(TunnelOutcome::Disconnected), Err(e) => { @@ -258,6 +276,8 @@ pub async fn connect_and_run( } } writer_result = &mut writer_handle => { + frame_tx.close(); + let _ = tokio::time::timeout(Duration::from_secs(1), &mut dispatch).await; match writer_result { Ok(()) => warn!("writer task exited normally, triggering reconnect"), Err(e) => { @@ -278,24 +298,37 @@ pub async fn connect_and_run( } _ = shutdown.changed() => { debug!("shutdown during tunnel dispatch"); + let _ = session_drain_tx.send(true); + let deadline = Duration::from_millis(state.config.tunnel_drain_deadline_ms).saturating_add(Duration::from_secs(1)); + let _ = tokio::time::timeout(deadline, &mut dispatch).await; Ok(TunnelOutcome::Shutdown) } + } }; // Drop our sender; the writer will exit once all stream handler clones // are also dropped (i.e. after they finish their in-flight work). drop(frame_tx); + forward_drain.abort(); + let _ = forward_drain.await; if !drain_signal.is_finished() { drain_signal.abort(); let _ = drain_signal.await; } - // Wait for the writer task to finish with a generous timeout — the - // dispatcher already waits up to 30s for stream handlers, so 35s here - // covers that plus a small margin. - // Skip if the writer already exited (the select branch that fired). if !writer_handle.is_finished() { - let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await; + let flush_timeout = if *session_drain_tx.borrow() { + Duration::from_millis(state.config.tunnel_drain_deadline_ms) + } else { + Duration::from_secs(1) + }; + if tokio::time::timeout(flush_timeout, &mut writer_handle) + .await + .is_err() + { + writer_handle.abort(); + let _ = writer_handle.await; + } } let connected_for = connected_at.elapsed(); diff --git a/apps/aether-tunnel/src/tunnel/dispatcher.rs b/apps/aether-tunnel/src/tunnel/dispatcher.rs index 73fb9d156..10b72e0d6 100644 --- a/apps/aether-tunnel/src/tunnel/dispatcher.rs +++ b/apps/aether-tunnel/src/tunnel/dispatcher.rs @@ -9,7 +9,7 @@ use std::time::Duration; use bytes::Bytes; use futures_util::StreamExt; use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore}; -use tokio::task::JoinHandle; +use tokio::task::{AbortHandle, JoinSet}; use tokio_tungstenite::tungstenite::Message; use tracing::{debug, error, info, warn}; @@ -41,13 +41,25 @@ impl AsRef<[u8]> for BudgetedFramePayload { enum StreamDispatchStatus { Delivered, Closed, - TimedOut, + Congested, } #[derive(Clone)] struct StreamDispatchTarget { body_tx: mpsc::Sender, response_window: Arc, + handler: Option, +} + +struct StreamCompletion { + stream_id: u32, + finished_tx: mpsc::UnboundedSender, +} + +impl Drop for StreamCompletion { + fn drop(&mut self) { + let _ = self.finished_tx.send(self.stream_id); + } } /// A request stream is identified by a non-zero id and may only be opened @@ -109,7 +121,7 @@ where // reopen the same id and bypass the stream admission limit. let mut active_handler_ids: HashSet = HashSet::new(); // Track spawned stream handlers so we can wait for them on shutdown - let mut handler_handles: Vec> = Vec::new(); + let mut handler_handles = JoinSet::new(); let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::(); let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize; let mut frames_since_cleanup: u32 = 0; @@ -121,30 +133,48 @@ where // Track last time we received any data to detect stale connections let mut last_data_at = tokio::time::Instant::now(); let mut draining = *drain.borrow(); + let mut drain_open = true; + let mut drain_deadline = draining.then(|| { + tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms) + }); + let mut initial_window_bytes = state.config.tunnel_stream_initial_window_bytes; + let mut close_rx = frame_tx.subscribe_close(); let read_err = loop { + if *close_rx.borrow() { + break None; + } if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after in-flight streams completed"); break None; } let msg_result = tokio::select! { + _ = close_rx.changed() => break None, msg = ws_stream.next() => { match msg { Some(r) => r, None => break None, } } - changed = drain.changed() => { + changed = drain.changed(), if drain_open => { if changed.is_err() { + drain_open = false; continue; } if *drain.borrow() { info!("tunnel drain requested, waiting for in-flight streams"); draining = true; + drain_deadline.get_or_insert_with(|| tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms)); } continue; } + _ = async { + match drain_deadline { + Some(deadline) => tokio::time::sleep_until(deadline).await, + None => std::future::pending().await, + } + } => break None, finished = handler_finished_rx.recv() => { if let Some(stream_id) = finished { active_handler_ids.remove(&stream_id); @@ -238,20 +268,7 @@ where continue; } if draining { - if frame_tx - .try_send(Frame::new( - frame.stream_id, - MsgType::StreamError, - 0, - Bytes::from("tunnel draining"), - )) - .is_err() - { - warn!( - stream_id = frame.stream_id, - "writer channel full, StreamError dropped during drain" - ); - } + try_send_stream_error(&frame_tx, frame.stream_id, "tunnel draining"); continue; } @@ -263,6 +280,11 @@ where Ok(p) => p, Err(e) => { warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed"); + try_send_stream_error( + &frame_tx, + frame.stream_id, + "invalid request metadata", + ); continue; } }; @@ -270,21 +292,11 @@ where Ok(m) => m, Err(e) => { warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata"); - // Use try_send to avoid blocking the read loop - if frame_tx - .try_send(Frame::new( - frame.stream_id, - MsgType::StreamError, - 0, - Bytes::from(format!("invalid request metadata: {e}")), - )) - .is_err() - { - warn!( - stream_id = frame.stream_id, - "writer channel full, StreamError dropped" - ); - } + try_send_stream_error( + &frame_tx, + frame.stream_id, + "invalid request metadata", + ); continue; } }; @@ -294,33 +306,27 @@ where stream_id = frame.stream_id, "max concurrent streams reached" ); - if frame_tx - .try_send(Frame::new( - frame.stream_id, - MsgType::StreamError, - 0, - Bytes::from("max concurrent streams reached"), - )) - .is_err() - { - warn!( - stream_id = frame.stream_id, - "writer channel full, StreamError dropped" - ); - } + try_send_stream_error( + &frame_tx, + frame.stream_id, + "max concurrent streams reached", + ); continue; } // Create body channel and spawn handler - let (body_tx, body_rx) = mpsc::channel::(64); - let response_window = Arc::new(StreamSendWindow::new( - state.config.tunnel_stream_initial_window_bytes, - )); + let body_capacity = (initial_window_bytes as usize) + .div_ceil(32 * 1024) + .saturating_add(1) + .max(64); + let (body_tx, body_rx) = mpsc::channel::(body_capacity); + let response_window = Arc::new(StreamSendWindow::new(initial_window_bytes)); streams.insert( frame.stream_id, StreamDispatchTarget { body_tx, response_window: Arc::clone(&response_window), + handler: None, }, ); active_handler_ids.insert(frame.stream_id); @@ -329,9 +335,13 @@ where let state_clone = Arc::clone(&state); let server_clone = Arc::clone(&server); let tx_clone = frame_tx.clone(); - let finished_tx = handler_finished_tx.clone(); let sid = frame.stream_id; - let handle = tokio::spawn(async move { + let completion = StreamCompletion { + stream_id: sid, + finished_tx: handler_finished_tx.clone(), + }; + let handle = handler_handles.spawn(async move { + let _completion = completion; stream_handler::handle_stream( state_clone, server_clone, @@ -342,9 +352,8 @@ where response_window, ) .await; - let _ = finished_tx.send(sid); }); - handler_handles.push(handle); + streams.get_mut(&sid).expect("new stream exists").handler = Some(handle); if request_headers_end_stream { if let Some(target) = streams.get(&sid) { @@ -365,19 +374,21 @@ where let is_end = frame.is_end_stream(); let sid = frame.stream_id; let dispatch = dispatch_stream_frame(&target.body_tx, frame).await; - if dispatch != StreamDispatchStatus::Delivered { - streams.remove(&sid); - if dispatch == StreamDispatchStatus::TimedOut { - server.tunnel_metrics.record_error( - "stream_dispatch_timeout", - &format!("request body dispatch timed out for stream {}", sid), - ); - try_send_stream_error( - &frame_tx, - sid, - "tunnel request body dispatch stalled", - ); + if dispatch == StreamDispatchStatus::Congested { + if let Some(target) = streams.remove(&sid) { + if let Some(handler) = target.handler { + handler.abort(); + } } + server.tunnel_metrics.record_error( + "stream_dispatch_timeout", + &format!("request body dispatch congested for stream {}", sid), + ); + try_send_stream_error( + &frame_tx, + sid, + "tunnel request body dispatch stalled", + ); if is_end && draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after request body completion"); @@ -387,10 +398,29 @@ where } } - MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => { + MsgType::StreamEnd => { + if let Some(target) = streams.get(&frame.stream_id) { + if dispatch_stream_frame(&target.body_tx, frame.clone()).await + == StreamDispatchStatus::Congested + { + if let Some(handler) = &target.handler { + handler.abort(); + } + try_send_stream_error( + &frame_tx, + frame.stream_id, + "tunnel request body dispatch stalled", + ); + } + } + } + + MsgType::StreamError | MsgType::ResetStream => { // Client-side cancellation or end if let Some(target) = streams.remove(&frame.stream_id) { - let _ = dispatch_stream_frame(&target.body_tx, frame).await; + if let Some(handler) = target.handler { + handler.abort(); + } if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after stream termination"); break None; @@ -409,12 +439,22 @@ where } MsgType::HeartbeatAck => { - heartbeat.on_ack(frame.payload).await; + heartbeat.on_ack(frame.payload); } MsgType::GoAway => { info!("received GOAWAY"); - break None; + draining = true; + let deadline_ms = + serde_json::from_slice::( + &frame.payload, + ) + .map(|payload| payload.drain_deadline_ms) + .unwrap_or(state.config.tunnel_drain_deadline_ms) + .min(state.config.tunnel_drain_deadline_ms); + drain_deadline.get_or_insert_with(|| { + tokio::time::Instant::now() + Duration::from_millis(deadline_ms) + }); } MsgType::WindowUpdate => { @@ -433,7 +473,31 @@ where ); } - MsgType::Hello | MsgType::Settings | MsgType::LoadReport => { + MsgType::Settings => { + if frame.stream_id != 0 || frame.flags != 0 { + break None; + } + let settings = serde_json::from_slice::( + &frame.payload, + ) + .ok() + .filter(|settings| settings.is_valid()); + let Some(settings) = settings else { + warn!("invalid tunnel SETTINGS"); + break None; + }; + if !streams.is_empty() + && settings.initial_stream_window_bytes != initial_window_bytes + { + warn!("tunnel SETTINGS changed with active streams"); + break None; + } + initial_window_bytes = settings + .initial_stream_window_bytes + .min(state.config.tunnel_stream_initial_window_bytes); + } + + MsgType::Hello | MsgType::LoadReport => { debug!( msg_type = ?frame.msg_type, stream_id = frame.stream_id, @@ -455,7 +519,7 @@ where // Trigger every 64 frames OR when the count exceeds max_streams. frames_since_cleanup += 1; if frames_since_cleanup >= 64 || handler_handles.len() > max_streams { - handler_handles.retain(|h| !h.is_finished()); + while handler_handles.try_join_next().is_some() {} frames_since_cleanup = 0; if draining && streams.is_empty() && active_handler_ids.is_empty() { info!("tunnel drained after cleanup"); @@ -467,9 +531,7 @@ where // Drop body senders so stream handlers waiting on body_rx will unblock streams.clear(); - // Wait for active stream handlers to finish so their frame_tx clones - // are dropped before the writer closes the sink. - drain_handlers(handler_handles).await; + handler_handles.shutdown().await; match read_err { Some(e) => Err(e.into()), @@ -478,30 +540,13 @@ where } async fn dispatch_stream_frame(tx: &mpsc::Sender, frame: Frame) -> StreamDispatchStatus { - let stream_id = frame.stream_id; - let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async { - let frame = attach_request_body_queue_budget(frame).await?; - tx.send(frame).await.ok()?; - Some(()) - }) - .await; - match dispatched { - Ok(Some(())) => StreamDispatchStatus::Delivered, - Ok(None) => { - warn!( - stream_id, - "stream handler channel or request body budget closed while dispatching tunnel frame" - ); - StreamDispatchStatus::Closed - } - Err(_) => { - warn!( - stream_id, - timeout_ms = stream_frame_dispatch_timeout().as_millis(), - "stream handler channel blocked while dispatching tunnel frame" - ); - StreamDispatchStatus::TimedOut - } + let Some(frame) = attach_request_body_queue_budget(frame).await else { + return StreamDispatchStatus::Congested; + }; + match tx.try_send(frame) { + Ok(()) => StreamDispatchStatus::Delivered, + Err(mpsc::error::TrySendError::Closed(_)) => StreamDispatchStatus::Closed, + Err(mpsc::error::TrySendError::Full(_)) => StreamDispatchStatus::Congested, } } @@ -523,7 +568,7 @@ async fn attach_request_body_queue_budget_with( return Some(frame); } let permits = request_body_queue_permits(&frame, budget_bytes)?; - let permit = budget.acquire_many_owned(permits).await.ok()?; + let permit = budget.try_acquire_many_owned(permits).ok()?; frame.payload = Bytes::from_owner(BudgetedFramePayload { bytes: frame.payload, _permit: permit, @@ -549,20 +594,6 @@ fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option u32::try_from(retained_bytes).ok() } -/// Bound how long a single stream handler is allowed to block the shared -/// WebSocket read loop while receiving request-body frames. -fn stream_frame_dispatch_timeout() -> Duration { - #[cfg(test)] - { - Duration::from_millis(25) - } - - #[cfg(not(test))] - { - Duration::from_millis(500) - } -} - fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) { if frame_tx .try_send(Frame::new( @@ -573,6 +604,7 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat )) .is_err() { + frame_tx.close(); warn!( stream_id, "writer channel full, StreamError dropped while aborting stalled stream" @@ -587,21 +619,6 @@ fn prune_closed_stream_senders(streams: &mut HashMap) before.saturating_sub(streams.len()) } -/// Wait for all active stream handlers to finish (with a timeout). -async fn drain_handlers(handles: Vec>) { - if handles.is_empty() { - return; - } - let count = handles.len(); - debug!(count, "waiting for active stream handlers to finish"); - let _ = tokio::time::timeout(Duration::from_secs(30), async { - for h in handles { - let _ = h.await; - } - }) - .await; -} - #[cfg(test)] mod tests { use super::*; @@ -633,7 +650,7 @@ mod tests { assert_eq!( stalled_send.await.expect("dispatch task should join"), - StreamDispatchStatus::TimedOut + StreamDispatchStatus::Congested ); let retained = rx @@ -737,6 +754,7 @@ mod tests { StreamDispatchTarget { body_tx: closed_tx, response_window: Arc::new(StreamSendWindow::new(1024)), + handler: None, }, ), ( @@ -744,6 +762,7 @@ mod tests { StreamDispatchTarget { body_tx: open_tx, response_window: Arc::new(StreamSendWindow::new(1024)), + handler: None, }, ), ]); @@ -763,6 +782,7 @@ mod tests { StreamDispatchTarget { body_tx: tx, response_window: Arc::new(StreamSendWindow::new(1024)), + handler: None, }, )]); let mut active_handler_ids = HashSet::from([7]); diff --git a/apps/aether-tunnel/src/tunnel/heartbeat.rs b/apps/aether-tunnel/src/tunnel/heartbeat.rs index 0a35cb09c..e6cc56409 100644 --- a/apps/aether-tunnel/src/tunnel/heartbeat.rs +++ b/apps/aether-tunnel/src/tunnel/heartbeat.rs @@ -31,14 +31,22 @@ enum AckDecision { } /// Handle for the dispatcher to forward HeartbeatAck frames. -#[derive(Clone)] pub struct HeartbeatHandle { ack_tx: tokio::sync::mpsc::Sender, + task: Option>, } impl HeartbeatHandle { - pub async fn on_ack(&self, payload: Bytes) { - let _ = self.ack_tx.send(payload).await; + pub fn on_ack(&self, payload: Bytes) { + let _ = self.ack_tx.try_send(payload); + } +} + +impl Drop for HeartbeatHandle { + fn drop(&mut self) { + if let Some(task) = self.task.take() { + task.abort(); + } } } @@ -48,7 +56,7 @@ impl HeartbeatHandle { pub fn spawn_noop() -> HeartbeatHandle { let (ack_tx, _) = tokio::sync::mpsc::channel::(1); // receiver is immediately dropped; on_ack() calls will silently fail - HeartbeatHandle { ack_tx } + HeartbeatHandle { ack_tx, task: None } } #[derive(Debug, Clone, Copy, Default)] @@ -74,7 +82,7 @@ pub fn spawn( ) -> HeartbeatHandle { let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::(4); - tokio::spawn(async move { + let task = tokio::spawn(async move { // Read initial interval from dynamic config (may be updated by remote config). let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval); let mut current_interval = initial_interval; @@ -151,7 +159,8 @@ pub fn spawn( current_interval = new_interval; } } - Some(ack_payload) = ack_rx.recv() => { + ack_payload = ack_rx.recv() => { + let Some(ack_payload) = ack_payload else { break; }; match handle_ack(&server, &ack_payload) { AckDecision::Accept { heartbeat_id: ack_id, @@ -179,7 +188,10 @@ pub fn spawn( } }); - HeartbeatHandle { ack_tx } + HeartbeatHandle { + ack_tx, + task: Some(task), + } } async fn build_heartbeat_payload( diff --git a/apps/aether-tunnel/src/tunnel/mod.rs b/apps/aether-tunnel/src/tunnel/mod.rs index 052b873b6..eca3d06e6 100644 --- a/apps/aether-tunnel/src/tunnel/mod.rs +++ b/apps/aether-tunnel/src/tunnel/mod.rs @@ -3,6 +3,7 @@ pub mod dispatcher; pub mod heartbeat; pub mod protocol; pub mod stream_handler; +mod task; pub mod writer; use std::sync::Arc; @@ -332,9 +333,9 @@ mod tests { ) .await; - assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1); - tokio::time::sleep(Duration::from_millis(200)).await; gateway_handle.abort(); + let _ = (&mut gateway_handle).await; + assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1); let (_restarted_gateway_state, restarted_gateway_handle) = start_gateway_on_port_retry(gateway_port) @@ -349,6 +350,7 @@ mod tests { ) .await; + assert!(server.tunnel_metrics.snapshot().connect_successes >= 2); let _ = shutdown_tx.send(true); tokio::time::timeout(Duration::from_secs(5), tunnel_task) .await @@ -380,7 +382,17 @@ mod tests { gateway_base_url: &str, node_id: &str, ) -> Option<(StatusCode, String)> { - let payload = relay_probe_envelope(); + let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?; + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + Some((status, body)) + } + + async fn relay_response( + gateway_base_url: &str, + node_id: &str, + payload: Vec, + ) -> Option { let timestamp = SystemTime::now() .duration_since(UNIX_EPOCH) .expect("test clock should be after epoch") @@ -398,7 +410,7 @@ mod tests { &nonce, &digest, ); - let response = reqwest::Client::new() + reqwest::Client::new() .post(format!( "{gateway_base_url}/api/internal/tunnel/relay/{node_id}" )) @@ -421,10 +433,7 @@ mod tests { .body(payload) .send() .await - .ok()?; - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Some((status, body)) + .ok() } fn relay_probe_envelope() -> Vec { @@ -456,30 +465,150 @@ mod tests { ) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> { // The embedded gateway now fails closed when relay authentication is // not configured. Keep this integration fixture explicitly authenticated. - let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET"); - let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID"); - std::env::set_var( - "AETHER_TUNNEL_RELAY_AUTH_SECRET", - "tunnel-reconnect-test-secret-at-least-32-bytes", - ); - std::env::set_var( - "AETHER_GATEWAY_INSTANCE_ID", - "tunnel-reconnect-test-gateway", - ); - let mut state = GatewayAppState::new().expect("gateway test state should build"); - aether_gateway::configure_test_tunnel_security( - &mut state, - "node-recovery", - "test-generation-1", - "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=", - ); - restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret); - restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance); + static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); + let state = { + let _guard = ENV_LOCK.lock().unwrap(); + let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET"); + let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID"); + std::env::set_var( + "AETHER_TUNNEL_RELAY_AUTH_SECRET", + "tunnel-reconnect-test-secret-at-least-32-bytes", + ); + std::env::set_var( + "AETHER_GATEWAY_INSTANCE_ID", + "tunnel-reconnect-test-gateway", + ); + let mut state = GatewayAppState::new().expect("gateway test state should build"); + aether_gateway::configure_test_tunnel_security( + &mut state, + "node-recovery", + "test-generation-1", + "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=", + ); + restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret); + restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance); + state + }; let router = build_router_with_state(state.clone()); let handle = spawn_router_on_port(port, router).await?; Ok((state, handle)) } + #[tokio::test] + async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() { + use axum::body::{Body, Bytes}; + use axum::routing::get; + use futures_util::StreamExt; + + ensure_rustls_provider(); + let upstream_port = reserve_local_port().unwrap(); + let upstream = Router::new() + .route( + "/large", + get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }), + ) + .route( + "/idle", + get(|| async { + let first = futures_util::stream::once(async { + Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n")) + }); + ( + [("content-type", "text/event-stream")], + Body::from_stream(first.chain(futures_util::stream::pending())), + ) + }), + ); + let upstream_task = super::task::SessionTask::new( + spawn_router_on_port(upstream_port, upstream).await.unwrap(), + ); + let gateway_port = reserve_local_port().unwrap(); + let gateway_url = format!("http://127.0.0.1:{gateway_port}"); + let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap(); + let gateway_task = super::task::SessionTask::new(gateway_task); + let mut config = sample_config(&gateway_url); + config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired; + config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into()); + config.tunnel_stream_initial_window_bytes = 512 * 1024; + config.tunnel_drain_deadline_ms = 100; + config.allow_private_targets = true; + config.allowed_ports.push(upstream_port); + let state = sample_state(config); + let server = sample_server(&state, "node-recovery"); + let (shutdown_tx, shutdown_rx) = watch::channel(false); + let (_drain_tx, drain_rx) = watch::channel(false); + let tunnel_task = super::task::SessionTask::new(tokio::spawn({ + let state = Arc::clone(&state); + let server = Arc::clone(&server); + async move { + run(&state, &server, 0, shutdown_rx, drain_rx).await; + } + })); + wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await; + + let envelope = |path: &str| { + let mut meta: protocol::RequestMeta = + serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap(); + meta.url = format!("http://127.0.0.1:{upstream_port}/{path}"); + meta.stream = true; + meta.timeout = 10; + meta.stream_first_byte_timeout_ms = Some(10_000); + let encoded = serde_json::to_vec(&meta).unwrap(); + let mut result = (encoded.len() as u32).to_be_bytes().to_vec(); + result.extend(encoded); + result + }; + let response = relay_response(&gateway_url, "node-recovery", envelope("large")) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let body = tokio::time::timeout(Duration::from_secs(10), response.bytes()) + .await + .unwrap() + .unwrap(); + assert_eq!(body.len(), 2 * 1024 * 1024); + assert!(body.iter().all(|byte| *byte == b'x')); + + let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) + .await + .unwrap(); + assert_eq!( + response.chunk().await.unwrap().unwrap(), + "data: started\n\n" + ); + drop(response); + tokio::time::timeout(Duration::from_secs(3), async { + while server + .active_connections + .load(std::sync::atomic::Ordering::Acquire) + != 0 + { + tokio::task::yield_now().await; + } + }) + .await + .expect("cancelled SSE must release the upstream handler"); + + let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle")) + .await + .unwrap(); + assert!(response.chunk().await.unwrap().is_some()); + shutdown_tx.send(true).unwrap(); + tokio::time::timeout(Duration::from_secs(3), tunnel_task) + .await + .unwrap() + .unwrap(); + assert_eq!( + server + .active_connections + .load(std::sync::atomic::Ordering::Acquire), + 0 + ); + drop(response); + drop(gateway_task); + drop(upstream_task); + } + fn restore_test_env(key: &str, value: Option) { if let Some(value) = value { std::env::set_var(key, value); diff --git a/apps/aether-tunnel/src/tunnel/stream_handler.rs b/apps/aether-tunnel/src/tunnel/stream_handler.rs index d77819713..39174fdc8 100644 --- a/apps/aether-tunnel/src/tunnel/stream_handler.rs +++ b/apps/aether-tunnel/src/tunnel/stream_handler.rs @@ -52,6 +52,7 @@ static REDIRECT_REPLAY_BUFFERED_BYTES: AtomicUsize = AtomicUsize::new(0); #[derive(Debug)] pub(crate) struct StreamSendWindow { + initial_window_bytes: u32, available: Mutex, notify: Notify, } @@ -59,6 +60,7 @@ pub(crate) struct StreamSendWindow { impl StreamSendWindow { pub(crate) fn new(initial_window_bytes: u32) -> Self { Self { + initial_window_bytes: initial_window_bytes.max(1), available: Mutex::new(u64::from(initial_window_bytes.max(1))), notify: Notify::new(), } @@ -82,6 +84,9 @@ impl StreamSendWindow { let requested = bytes as u64; let started_at = Instant::now(); loop { + let notified = self.notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); { let mut available = self.available.lock().expect("stream window lock poisoned"); if *available >= requested { @@ -93,10 +98,7 @@ impl StreamSendWindow { let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else { return Err(()); }; - if tokio::time::timeout(remaining, self.notify.notified()) - .await - .is_err() - { + if tokio::time::timeout(remaining, notified).await.is_err() { return Err(()); } } @@ -173,31 +175,33 @@ fn safe_stream_error_message(message: &str) -> &'static str { "upstream request failed" } -fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) { +async fn send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) -> bool { if bytes == 0 { - return; + return true; } let delta = bytes.min(u32::MAX as usize) as u32; - if frame_tx - .try_send(TunnelFrame::new( - stream_id, - MsgType::WindowUpdate, - 0, - Bytes::from( - serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload { - delta_bytes: delta, - }) - .expect("window update payload should serialize"), - ), - )) - .is_err() - { - warn!( - stream_id, - delta_bytes = delta, - "writer channel full, WINDOW_UPDATE dropped" - ); + if matches!( + tokio::time::timeout( + FLOW_CONTROL_WAIT_TIMEOUT, + frame_tx.send(TunnelFrame::new( + stream_id, + MsgType::WindowUpdate, + 0, + Bytes::from( + serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload { + delta_bytes: delta, + }) + .expect("window update payload should serialize"), + ), + )) + ) + .await, + Ok(Ok(())) + ) { + return true; } + frame_tx.close(); + false } /// Match reqwest's default redirect budget so direct execution and tunnel relay @@ -242,6 +246,23 @@ enum ReplayableRequestBody { struct PreparedRequestBody { first_request_body: Option, replay_body: ReplayableRequestBody, + spool_task: Option>, +} + +impl Drop for PreparedRequestBody { + fn drop(&mut self) { + if let Some(task) = self.spool_task.take() { + task.abort(); + } + } +} + +struct ActiveStreamGuard(Arc); + +impl Drop for ActiveStreamGuard { + fn drop(&mut self) { + self.0.active_connections.fetch_sub(1, Ordering::Release); + } } #[derive(Debug, Clone, Copy)] @@ -331,7 +352,10 @@ impl hyper::body::Body for ReplayRequestBody { #[derive(Debug)] enum SpoolBodyEvent { - Data(Bytes), + Data { + payload: Bytes, + credit_returned: bool, + }, Error(String), End, } @@ -563,8 +587,9 @@ impl RequestBodyReplayState { } } - fn push_chunk(&self, payload: Bytes) { + fn push_chunk(&self, payload: Bytes) -> bool { let mut disable_replay = false; + let mut retained = false; let mut state = self.state.lock().expect("request body replay state lock"); if let RequestBodyReplayStatus::Collecting { chunks, @@ -577,7 +602,7 @@ impl RequestBodyReplayState { drop(state); self.release_reserved_bytes(); self.ready.notify_waiters(); - return; + return false; }; let accounted_bytes = payload.len().checked_add(std::mem::size_of::()); if next_len > self.budget_bytes @@ -590,6 +615,7 @@ impl RequestBodyReplayState { } else { *buffered_len = next_len; chunks.push(payload); + retained = true; } } drop(state); @@ -597,6 +623,7 @@ impl RequestBodyReplayState { self.release_reserved_bytes(); self.ready.notify_waiters(); } + retained } fn try_reserve_bytes(&self, bytes: usize) -> bool { @@ -951,9 +978,6 @@ pub(super) fn decode_request_body_frame(frame: TunnelFrame) -> Result, @@ -973,19 +997,20 @@ fn prepare_request_body( None => ReplayableRequestBody::NonReplayable, }; - tokio::spawn(spool_request_body( + let spool_task = tokio::spawn(spool_request_body( stream_id, body_rx, spool_tx, replay_state, body_size, deadline, - frame_tx, + frame_tx.clone(), )); PreparedRequestBody { - first_request_body: Some(build_spooled_request_body(spool_rx)), + first_request_body: Some(build_spooled_request_body(spool_rx, stream_id, frame_tx)), replay_body, + spool_task: Some(spool_task), } } @@ -1001,6 +1026,7 @@ fn prepare_bodyless_request_body( } else { ReplayableRequestBody::NonReplayable }, + spool_task: None, } } @@ -1056,11 +1082,16 @@ async fn spool_request_body( }; let Some(frame) = frame else { + let message = "tunnel request body closed before stream end".to_string(); if let Some(state) = &replay_state { - state.finish(); + state.fail(message.clone()); } - let _ = - send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()).await; + let _ = send_spool_event( + &mut spool_tx, + SpoolBodyEvent::Error(message), + replay_state.as_ref(), + ) + .await; return; }; @@ -1086,13 +1117,23 @@ async fn spool_request_body( if !payload.is_empty() { body_size.fetch_add(payload.len(), Ordering::Relaxed); - try_send_window_update(&frame_tx, stream_id, payload.len()); - if let Some(state) = &replay_state { - state.push_chunk(payload.clone()); + let credit_returned = replay_state + .as_ref() + .is_some_and(|state| state.push_chunk(payload.clone())); + if credit_returned + && !send_window_update(&frame_tx, stream_id, payload.len()).await + { + if let Some(state) = &replay_state { + state.fail("tunnel flow-control update failed".to_string()); + } + return; } if send_spool_event( &mut spool_tx, - SpoolBodyEvent::Data(payload), + SpoolBodyEvent::Data { + payload, + credit_returned, + }, replay_state.as_ref(), ) .await @@ -1479,6 +1520,7 @@ where } let mut stream = response.into_body().into_data_stream(); + let chunk_size = MAX_CHUNK_SIZE.min(response_window.initial_window_bytes as usize); loop { let chunk_result = if let Some(deadline) = response_body_deadline { let Some(remaining) = remaining_timeout(deadline) else { @@ -1531,7 +1573,7 @@ where match chunk_result { Ok(chunk) => { - if chunk.len() <= MAX_CHUNK_SIZE { + if chunk.len() <= chunk_size { let (payload, extra_flags) = raw_payload(chunk); if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len()) .await @@ -1561,7 +1603,7 @@ where } else { let mut offset = 0; while offset < chunk.len() { - let end = (offset + MAX_CHUNK_SIZE).min(chunk.len()); + let end = (offset + chunk_size).min(chunk.len()); let slice = chunk.slice(offset..end); let (payload, extra_flags) = raw_payload(slice); if !acquire_response_credit( @@ -1735,6 +1777,7 @@ pub async fn handle_stream( }; server.active_connections.fetch_add(1, Ordering::Release); + let _active_stream = ActiveStreamGuard(Arc::clone(&server)); let stream_io = StreamIo { body_rx, @@ -1745,7 +1788,6 @@ pub async fn handle_stream( let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await; - server.active_connections.fetch_sub(1, Ordering::Release); if let Some(d) = connect_elapsed { server.metrics.record_request(d); } @@ -1772,6 +1814,18 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool { timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64, "writer channel stalled for body frame, abandoning stream" ); + let reset = TunnelFrame::new( + stream_id, + MsgType::ResetStream, + 0, + Bytes::from_static(b"{\"reason\":\"tunnel writer stalled\"}"), + ); + if !matches!( + tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(reset)).await, + Ok(Ok(())) + ) { + tx.close(); + } false } Ok(Err(QueueSendError::Full(_))) => { @@ -1781,7 +1835,10 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool { } else { match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await { Ok(Ok(())) => true, - Ok(Err(_)) => false, + Ok(Err(_)) => { + tx.close(); + false + } Err(_) => { warn!( stream_id, @@ -1789,6 +1846,7 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool { flags = flags, "control frame send timeout (writer congested), abandoning stream" ); + tx.close(); false } } @@ -2126,7 +2184,6 @@ async fn handle_stream_inner( } async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) { - // Error frames use best-effort delivery — don't block if writer is congested let safe_message = safe_stream_error_message(msg); let _ = send_frame( tx, @@ -2163,22 +2220,42 @@ fn build_streaming_request_body( fn build_spooled_request_body( spool_rx: mpsc::Receiver, + stream_id: u32, + frame_tx: FrameSender, ) -> upstream_client::UpstreamRequestBody { - let body_stream = stream::unfold((spool_rx, false), |(mut spool_rx, finished)| async move { - if finished { - return None; - } + let body_stream = stream::unfold( + (spool_rx, frame_tx, false), + move |(mut spool_rx, frame_tx, finished)| async move { + if finished { + return None; + } - match spool_rx.recv().await { - Some(SpoolBodyEvent::Data(payload)) => { - Some((Ok(BodyFrame::data(payload)), (spool_rx, false))) + match spool_rx.recv().await { + Some(SpoolBodyEvent::Data { + payload, + credit_returned, + }) => { + if !credit_returned + && !send_window_update(&frame_tx, stream_id, payload.len()).await + { + return Some(( + Err(io::Error::other("tunnel flow-control update failed")), + (spool_rx, frame_tx, true), + )); + } + Some((Ok(BodyFrame::data(payload)), (spool_rx, frame_tx, false))) + } + Some(SpoolBodyEvent::Error(message)) => { + Some((Err(io::Error::other(message)), (spool_rx, frame_tx, true))) + } + Some(SpoolBodyEvent::End) => None, + None => Some(( + Err(io::Error::other("tunnel request body ended unexpectedly")), + (spool_rx, frame_tx, true), + )), } - Some(SpoolBodyEvent::Error(message)) => { - Some((Err(io::Error::other(message)), (spool_rx, true))) - } - Some(SpoolBodyEvent::End) | None => None, - } - }); + }, + ); upstream_client::stream_request_body(body_stream) } @@ -2249,6 +2326,105 @@ fn build_prefixed_request_body( #[cfg(test)] mod tests { + #[tokio::test(start_paused = true)] + async fn window_updates_wait_for_capacity_instead_of_disappearing() { + let (high_tx, mut high_rx) = aether_runtime::bounded_queue(1); + let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1); + let sender = FrameSender::from_test_queues(high_tx, normal_tx); + sender + .try_send(TunnelFrame::control(MsgType::Ping, Bytes::new())) + .unwrap(); + let task = tokio::spawn(async move { send_window_update(&sender, 7, 1024).await }); + tokio::time::sleep(Duration::from_secs(1)).await; + assert!(!task.is_finished()); + high_rx.recv().await.unwrap(); + assert!(task.await.unwrap()); + let update = high_rx.recv().await.unwrap(); + assert_eq!(update.msg_type, MsgType::WindowUpdate); + let payload: aether_contracts::tunnel::WindowUpdatePayload = + serde_json::from_slice(&update.payload).unwrap(); + assert_eq!(payload.delta_bytes, 1024); + } + + #[tokio::test(start_paused = true)] + async fn stalled_body_delivery_emits_a_reset() { + let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4); + let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1); + let sender = FrameSender::from_test_queues(high_tx, normal_tx); + sender + .try_send(TunnelFrame::new( + 7, + MsgType::ResponseBody, + 0, + Bytes::from_static(b"first"), + )) + .unwrap(); + assert!( + !send_frame( + &sender, + TunnelFrame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"second")) + ) + .await + ); + let reset = high_rx.recv().await.unwrap(); + assert_eq!(reset.msg_type, MsgType::ResetStream); + assert_eq!(reset.stream_id, 7); + } + + #[tokio::test] + async fn request_credit_follows_consumption_without_redirect_replay() { + let (body_tx, body_rx) = mpsc::channel(4); + let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4); + let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4); + let sender = FrameSender::from_test_queues(high_tx, normal_tx); + let mut prepared = prepare_request_body( + 7, + body_rx, + Arc::new(AtomicUsize::new(0)), + Instant::now() + Duration::from_secs(10), + false, + sender, + ); + body_tx + .send(TunnelFrame::new( + 7, + MsgType::RequestBody, + flags::END_STREAM, + Bytes::from_static(b"body"), + )) + .await + .unwrap(); + tokio::task::yield_now().await; + assert!(high_rx.try_recv().is_err()); + let mut body = prepared.take_first_request_body(); + assert!(body.frame().await.unwrap().is_ok()); + assert_eq!( + high_rx.recv().await.unwrap().msg_type, + MsgType::WindowUpdate + ); + assert!(body.frame().await.is_none()); + } + + #[tokio::test] + async fn dropping_prepared_body_cancels_its_spooler() { + let (body_tx, body_rx) = mpsc::channel(4); + let (high_tx, _high_rx) = aether_runtime::bounded_queue(4); + let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4); + let sender = FrameSender::from_test_queues(high_tx, normal_tx); + let prepared = prepare_request_body( + 7, + body_rx, + Arc::new(AtomicUsize::new(0)), + Instant::now() + Duration::from_secs(3600), + false, + sender, + ); + drop(prepared); + tokio::time::timeout(Duration::from_secs(1), body_tx.closed()) + .await + .unwrap(); + } + use std::collections::HashMap; use std::net::SocketAddr; use std::pin::Pin; @@ -2379,7 +2555,7 @@ mod tests { let (tx, rx) = mpsc::channel(4); let (frame_tx, sent, writer_handle) = spawn_test_writer(); let body_size = Arc::new(AtomicUsize::new(0)); - let prepared = prepare_request_body( + let mut prepared = prepare_request_body( 1, rx, Arc::clone(&body_size), @@ -2389,6 +2565,7 @@ mod tests { ); let mut body = prepared .first_request_body + .take() .expect("first request body should be present"); tx.send(TunnelFrame::new( diff --git a/apps/aether-tunnel/src/tunnel/task.rs b/apps/aether-tunnel/src/tunnel/task.rs new file mode 100644 index 000000000..61f2e3d08 --- /dev/null +++ b/apps/aether-tunnel/src/tunnel/task.rs @@ -0,0 +1,52 @@ +use std::future::Future; +use std::pin::Pin; +use std::task::{Context, Poll}; + +use tokio::task::{JoinError, JoinHandle}; + +pub(super) struct SessionTask(JoinHandle); + +impl SessionTask { + pub(super) fn new(handle: JoinHandle) -> Self { + Self(handle) + } + pub(super) fn abort(&self) { + self.0.abort(); + } + pub(super) fn is_finished(&self) -> bool { + self.0.is_finished() + } +} + +impl Future for SessionTask { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + Pin::new(&mut self.0).poll(context) + } +} + +impl Drop for SessionTask { + fn drop(&mut self) { + self.0.abort(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn dropping_a_session_task_aborts_its_child() { + let child = tokio::spawn(std::future::pending::<()>()); + let abort = child.abort_handle(); + drop(SessionTask::new(child)); + tokio::time::timeout(std::time::Duration::from_secs(1), async { + while !abort.is_finished() { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + } +} diff --git a/apps/aether-tunnel/src/tunnel/writer.rs b/apps/aether-tunnel/src/tunnel/writer.rs index f9929f0bc..5f22d0ca8 100644 --- a/apps/aether-tunnel/src/tunnel/writer.rs +++ b/apps/aether-tunnel/src/tunnel/writer.rs @@ -13,6 +13,7 @@ use aether_contracts::tunnel::{MsgType, HEADER_SIZE}; use aether_runtime::QueueSnapshot; use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError}; use futures_util::SinkExt; +use tokio::sync::watch; use tokio::task::JoinHandle; use tokio_tungstenite::tungstenite::Message; use tracing::{debug, error, trace}; @@ -24,6 +25,8 @@ use aether_contracts::tunnel_security::SecureFrameCodec; const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64; const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256; +const WRITE_TIMEOUT: Duration = Duration::from_secs(15); +const CLOSE_TIMEOUT: Duration = Duration::from_secs(1); #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum FramePriority { @@ -43,9 +46,18 @@ pub struct FrameQueueSnapshots { pub struct FrameSender { high_tx: BoundedQueueSender, normal_tx: BoundedQueueSender, + close_tx: watch::Sender, } impl FrameSender { + pub fn close(&self) { + let _ = self.close_tx.send(true); + } + + pub(super) fn subscribe_close(&self) -> watch::Receiver { + self.close_tx.subscribe() + } + pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError> { match classify_frame_priority(&frame) { FramePriority::High => self.high_tx.send(frame).await, @@ -73,7 +85,12 @@ impl FrameSender { high_tx: BoundedQueueSender, normal_tx: BoundedQueueSender, ) -> Self { - Self { high_tx, normal_tx } + let (close_tx, _) = watch::channel(false); + Self { + high_tx, + normal_tx, + close_tx, + } } } @@ -113,15 +130,24 @@ where { let (high_tx, mut high_rx) = bounded_queue::(HIGH_PRIORITY_QUEUE_CAPACITY); let (normal_tx, mut normal_rx) = bounded_queue::(NORMAL_PRIORITY_QUEUE_CAPACITY); - let tx = FrameSender { high_tx, normal_tx }; + let (close_tx, mut close_rx) = watch::channel(false); + let tx = FrameSender { + high_tx, + normal_tx, + close_tx, + }; let handle = tokio::spawn(async move { let mut ping_ticker = tokio::time::interval(ping_interval); let mut high_open = true; let mut normal_open = true; + let mut close_open = true; ping_ticker.tick().await; // skip first immediate tick loop { + if *close_rx.borrow() { + break; + } if let Ok(frame) = high_rx.try_recv() { if !write_frame( &mut sink, @@ -141,6 +167,10 @@ where tokio::select! { biased; + changed = close_rx.changed(), if close_open => { + if changed.is_err() { close_open = false; } + if *close_rx.borrow() { break; } + }, frame = high_rx.recv(), if high_open => { match frame { Some(frame) => { @@ -152,7 +182,7 @@ where } } _ = ping_ticker.tick(), if high_open || normal_open => { - if let Err(e) = sink.send(Message::Ping(vec![])).await { + if let Err(e) = send_message(&mut sink, Message::Ping(vec![])).await { error!(error = %e, "failed to send WebSocket ping"); if let Some(metrics) = tunnel_metrics.as_deref() { metrics.record_error("ws_ping_error", &e.to_string()); @@ -174,7 +204,7 @@ where } } debug!("writer task exiting"); - let _ = sink.close().await; + let _ = tokio::time::timeout(CLOSE_TIMEOUT, sink.close()).await; }); (tx, handle) @@ -228,7 +258,7 @@ where None => frame.encode(), }; let wire_len = data.len().max(HEADER_SIZE); - if let Err(e) = sink.send(Message::Binary(data.into())).await { + if let Err(e) = send_message(sink, Message::Binary(data.into())).await { error!( stream_id = stream_id, msg_type = ?msg_type, @@ -248,8 +278,92 @@ where true } +async fn send_message( + sink: &mut S, + message: Message, +) -> Result<(), tokio_tungstenite::tungstenite::Error> +where + S: SinkExt + Unpin, +{ + tokio::time::timeout(WRITE_TIMEOUT, sink.send(message)) + .await + .map_err(|_| { + tokio_tungstenite::tungstenite::Error::Io(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "tunnel WebSocket write timed out", + )) + })? +} + #[cfg(test)] mod tests { + #[tokio::test] + async fn dropping_last_sender_flushes_queued_body_and_end_frames() { + let sink = VecSink::default(); + let sent = Arc::clone(&sink.sent); + let (sender, task) = spawn_writer(sink, Duration::from_secs(60)); + sender + .send(Frame::new( + 7, + MsgType::ResponseBody, + 0, + bytes::Bytes::from_static(b"late"), + )) + .await + .unwrap(); + sender + .send(Frame::new(7, MsgType::StreamEnd, 0, bytes::Bytes::new())) + .await + .unwrap(); + drop(sender); + task.await.unwrap(); + let frames = sent.lock().unwrap(); + assert_eq!(frames.len(), 2); + let Message::Binary(body) = &frames[0] else { + panic!("expected body") + }; + assert_eq!( + Frame::decode(body.clone().into()).unwrap().payload, + b"late".as_slice() + ); + } + + struct StalledSink; + + impl futures_util::Sink for StalledSink { + type Error = Error; + fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Pending + } + fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Error> { + Ok(()) + } + fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Pending + } + fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Pending + } + } + + #[tokio::test(start_paused = true)] + async fn stalled_socket_write_and_close_are_bounded() { + let (sender, task) = spawn_writer(StalledSink, Duration::from_secs(60)); + sender + .send(Frame::new( + 1, + MsgType::ResponseBody, + 0, + bytes::Bytes::from_static(b"data"), + )) + .await + .unwrap(); + tokio::time::timeout(Duration::from_secs(20), task) + .await + .expect("writer should time out") + .unwrap(); + } + use std::pin::Pin; use std::sync::{Arc, Mutex}; use std::task::{Context, Poll}; diff --git a/crates/aether-contracts/src/tunnel.rs b/crates/aether-contracts/src/tunnel.rs index 115454b11..c07a9bda3 100644 --- a/crates/aether-contracts/src/tunnel.rs +++ b/crates/aether-contracts/src/tunnel.rs @@ -588,6 +588,68 @@ pub struct SettingsPayload { pub drain_deadline_ms: u64, } +impl SettingsPayload { + pub fn is_valid(&self) -> bool { + self.initial_stream_window_bytes > 0 + && u64::from(self.initial_stream_window_bytes) + <= MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64 + && self.min_window_update_bytes > 0 + && self.min_window_update_bytes <= self.initial_stream_window_bytes + && self.drain_deadline_ms > 0 + } + + pub fn negotiate(&self, initial_window_bytes: u32, drain_deadline_ms: u64) -> Self { + let window = self + .initial_stream_window_bytes + .min(initial_window_bytes) + .max(1); + Self { + initial_stream_window_bytes: window, + min_window_update_bytes: self.min_window_update_bytes.min((window / 4).max(1)), + drain_deadline_ms: self.drain_deadline_ms.min(drain_deadline_ms), + } + } +} + +#[cfg(test)] +mod settings_tests { + use super::*; + + #[test] + fn negotiation_bounds_window_updates_by_the_smaller_window() { + let settings = SettingsPayload { + initial_stream_window_bytes: 512 * 1024, + min_window_update_bytes: 128 * 1024, + drain_deadline_ms: 30_000, + }; + let negotiated = settings.negotiate(4 * 1024 * 1024, 1000); + assert_eq!(negotiated.initial_stream_window_bytes, 512 * 1024); + assert_eq!(negotiated.min_window_update_bytes, 128 * 1024); + assert_eq!(negotiated.drain_deadline_ms, 1000); + assert!(negotiated.is_valid()); + let tiny = settings.negotiate(1, 1); + assert_eq!(tiny.initial_stream_window_bytes, 1); + assert_eq!(tiny.min_window_update_bytes, 1); + assert!(tiny.is_valid()); + } + + #[test] + fn invalid_window_settings_are_rejected() { + let mut settings = SettingsPayload { + initial_stream_window_bytes: 1024, + min_window_update_bytes: 256, + drain_deadline_ms: 1, + }; + settings.min_window_update_bytes = 1025; + assert!(!settings.is_valid()); + settings.min_window_update_bytes = 0; + assert!(!settings.is_valid()); + settings.initial_stream_window_bytes = u32::MAX; + settings.min_window_update_bytes = 1; + assert!(!settings.is_valid()); + } +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub struct WindowUpdatePayload { pub delta_bytes: u32,