mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
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:
@@ -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<u64>,
|
||||
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<upstream_client::UpstreamRequestBody>,
|
||||
replay_body: ReplayableRequestBody,
|
||||
spool_task: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl Drop for PreparedRequestBody {
|
||||
fn drop(&mut self) {
|
||||
if let Some(task) = self.spool_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct ActiveStreamGuard(Arc<ServerContext>);
|
||||
|
||||
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::<Bytes>());
|
||||
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<Bytes, std
|
||||
Ok(frame.payload)
|
||||
}
|
||||
|
||||
// Drain tunnel body frames on a detached task so the shared dispatcher is no
|
||||
// longer coupled to upstream body polling. Redirect replay retains a bounded
|
||||
// copy; crossing either replay budget only disables replay for this request.
|
||||
fn prepare_request_body(
|
||||
stream_id: u32,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
@@ -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<SpoolBodyEvent>,
|
||||
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(
|
||||
|
||||
Reference in New Issue
Block a user