fix(tunnel): prevent stream stalls and harden session cleanup

Reliably deliver flow-control credits and terminal states, isolate slow streams and heartbeats, negotiate stream windows, and clean up cancelled streams and session tasks.

Add regression coverage for queue pressure, early cancellation, small-window streaming, drain, and reconnect. Validate 185 agent tests, 88 gateway tunnel tests, and 21 protocol tests.
This commit is contained in:
elky
2026-09-07 15:39:40 +08:00
parent aa7dbe67d3
commit ec95f2ca1f
17 changed files with 1883 additions and 354 deletions
@@ -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<BytesMut>,
bytes: usize,
terminal: Option<Result<(), String>>,
receiver_taken: bool,
receiver_closed: bool,
}
pub(super) struct ResponseBuffer {
state: Mutex<BufferState>,
notify: Notify,
capacity: usize,
}
impl ResponseBuffer {
pub(super) fn new(capacity: usize) -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(BufferState::default()),
notify: Notify::new(),
capacity,
})
}
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
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<ResponseBuffer>,
finished: bool,
}
impl BodyReceiver {
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
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")
);
}
}
@@ -0,0 +1,250 @@
use super::*;
async fn fixture(
window: u32,
capacity: usize,
) -> (
Arc<HubRouter>,
Arc<ProxyConn>,
Arc<LocalStream>,
aether_runtime::BoundedQueueReceiver<Message>,
) {
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<HubRouter>, 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();
}
}
+203 -83
View File
@@ -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<u32> = LazyLock::new(|| {
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
.ok()
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
});
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = 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<u64>,
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<bool> {
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<protocol::SettingsPayload>,
}
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<LocalResponseHead>,
@@ -585,10 +609,11 @@ pub struct LocalStream {
proxy_stream_id: u32,
request_window: StreamFlowWindow,
response_consumed_since_update: Mutex<u64>,
min_window_update_bytes: u32,
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
wait_state: Mutex<LocalWaitState>,
headers_notify: Notify,
body_tx: mpsc::Sender<LocalBodyEvent>,
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
body: Arc<ResponseBuffer>,
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<u32> {
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<LocalResponseHead, String> {
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<mpsc::Receiver<LocalBodyEvent>> {
self.body_rx.lock().take()
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
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<String>) {
@@ -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<LocalStream>,
failed: bool,
pending: Option<LocalBodyEvent>,
}
impl LocalBodyReceiver {
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
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<HashMap<String, u64>>,
}
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<String>,
@@ -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<Self>, 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::<protocol::SettingsPayload>(&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()));
}
}
}
}
@@ -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<LocalBodyEvent>,
body_rx: LocalBodyReceiver,
request_guard: StreamGuard,
_request_permit: Option<AdmissionPermit>,
}
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
}
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, 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 {
@@ -1,3 +1,4 @@
mod body;
mod control_plane;
mod hub;
mod local_relay;
@@ -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::<Message>(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<WebSocket, Message>,
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
security: Option<&SecureFrameCodec>,
protocol_version: u8,
) -> Option<protocol::SettingsPayload> {
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::<HelloPayload>(&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::<protocol::SettingsPayload>(&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::<std::net::SocketAddr>(),
)
.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=";