Refactor tunnel stability protocol

This commit is contained in:
elky
2026-06-01 01:36:49 +08:00
parent 392353ffff
commit 37413c0211
21 changed files with 1960 additions and 132 deletions
+53 -3
View File
@@ -18,7 +18,8 @@ use crate::egress_proxy::{
};
use crate::state::{AppState, ServerContext};
use aether_contracts::tunnel::{
CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_NODE_NAME_B64_HEADER, TUNNEL_PROTOCOL_VERSION_HEADER,
HelloPayload, SettingsPayload, CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_NODE_NAME_B64_HEADER,
TUNNEL_PROTOCOL_VERSION_HEADER,
};
use aether_contracts::tunnel_security::{
SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED,
@@ -186,7 +187,13 @@ pub async fn connect_and_run(
Some(Arc::clone(&server.tunnel_metrics)),
security.clone(),
);
let drain_signal = spawn_drain_signal(conn_idx, frame_tx.clone(), drain.clone());
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
let drain_signal = spawn_drain_signal(
conn_idx,
frame_tx.clone(),
drain.clone(),
state.config.tunnel_drain_deadline_ms,
);
// Spawn heartbeat task (only for primary connection to avoid
// resetting shared atomic metrics via swap(0))
@@ -298,10 +305,48 @@ pub async fn connect_and_run(
outcome
}
async fn send_protocol_v3_hello(
frame_tx: &writer::FrameSender,
security_session: &str,
state: &Arc<AppState>,
) {
let hello = super::protocol::Frame::control(
super::protocol::MsgType::Hello,
serde_json::to_vec(&HelloPayload {
protocol_version: CURRENT_TUNNEL_PROTOCOL_VERSION,
capabilities: vec![
"flow-control".to_string(),
"reset-stream".to_string(),
"graceful-drain".to_string(),
"load-report".to_string(),
],
session_id: Some(security_session.to_string()),
replica_id: None,
})
.expect("hello payload should serialize"),
);
let settings = super::protocol::Frame::control(
super::protocol::MsgType::Settings,
serde_json::to_vec(&SettingsPayload {
initial_stream_window_bytes: state.config.tunnel_stream_initial_window_bytes,
min_window_update_bytes: state
.config
.tunnel_stream_initial_window_bytes
.saturating_div(4)
.max(1),
drain_deadline_ms: state.config.tunnel_drain_deadline_ms,
})
.expect("settings payload should serialize"),
);
let _ = tokio::time::timeout(Duration::from_millis(250), frame_tx.send(hello)).await;
let _ = tokio::time::timeout(Duration::from_millis(250), frame_tx.send(settings)).await;
}
fn spawn_drain_signal(
conn_idx: usize,
frame_tx: writer::FrameSender,
mut drain: watch::Receiver<bool>,
drain_deadline_ms: u64,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
if !*drain.borrow() {
@@ -320,7 +365,12 @@ fn spawn_drain_signal(
Duration::from_millis(250),
frame_tx.send(super::protocol::Frame::control(
super::protocol::MsgType::GoAway,
bytes::Bytes::new(),
serde_json::to_vec(&aether_contracts::tunnel::GoAwayPayload {
last_accepted_stream_id: u32::MAX,
drain_deadline_ms,
reason: "tunnel drain requested".to_string(),
})
.expect("goaway payload should serialize"),
)),
)
.await
+99 -19
View File
@@ -17,6 +17,7 @@ use crate::state::{AppState, ServerContext};
use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::stream_handler;
use super::stream_handler::StreamSendWindow;
use super::writer::FrameSender;
use aether_contracts::tunnel_security::SecureFrameCodec;
@@ -27,6 +28,12 @@ enum StreamDispatchStatus {
TimedOut,
}
#[derive(Clone)]
struct StreamDispatchTarget {
body_tx: mpsc::Sender<Frame>,
response_window: Arc<StreamSendWindow>,
}
/// Run the dispatcher loop, reading from the WebSocket stream.
#[allow(dead_code)]
pub async fn run<S>(
@@ -61,10 +68,11 @@ where
+ Send
+ 'static,
{
// Active streams: stream_id -> body sender
let mut streams: HashMap<u32, mpsc::Sender<Frame>> = HashMap::new();
// Active streams: stream_id -> body sender + response flow-control window.
let mut streams: HashMap<u32, StreamDispatchTarget> = HashMap::new();
// Track spawned stream handlers so we can wait for them on shutdown
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
let mut frames_since_cleanup: u32 = 0;
let stale_timeout = state
@@ -99,6 +107,16 @@ where
}
continue;
}
finished = handler_finished_rx.recv() => {
if let Some(stream_id) = finished {
streams.remove(&stream_id);
if draining && streams.is_empty() {
info!("tunnel drained after stream handler completion");
break None;
}
}
continue;
}
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
warn!(
stale_ms = stale_timeout.as_millis(),
@@ -239,14 +257,22 @@ where
// Create body channel and spawn handler
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
let response_window = Arc::new(StreamSendWindow::new(
state.config.tunnel_stream_initial_window_bytes,
));
streams.insert(
frame.stream_id,
StreamDispatchTarget {
body_tx,
response_window: Arc::clone(&response_window),
},
);
let request_headers_end_stream = frame.is_end_stream();
if !request_headers_end_stream {
streams.insert(frame.stream_id, body_tx);
}
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 {
stream_handler::handle_stream(
@@ -256,20 +282,33 @@ where
meta,
body_rx,
tx_clone,
response_window,
)
.await;
let _ = finished_tx.send(sid);
});
handler_handles.push(handle);
if request_headers_end_stream {
if let Some(target) = streams.get(&sid) {
let _ = target.body_tx.try_send(Frame::new(
sid,
MsgType::StreamEnd,
0,
Bytes::new(),
));
}
}
debug!(stream_id = frame.stream_id, "new stream started");
}
MsgType::RequestBody => {
if let Some(tx) = streams.get(&frame.stream_id).cloned() {
if let Some(target) = streams.get(&frame.stream_id).cloned() {
let is_end = frame.is_end_stream();
let sid = frame.stream_id;
let dispatch = dispatch_stream_frame(&tx, frame).await;
if is_end || dispatch != StreamDispatchStatus::Delivered {
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(
@@ -282,7 +321,7 @@ where
"tunnel request body dispatch stalled",
);
}
if draining && streams.is_empty() {
if is_end && draining && streams.is_empty() {
info!("tunnel drained after request body completion");
break None;
}
@@ -290,10 +329,10 @@ where
}
}
MsgType::StreamEnd | MsgType::StreamError => {
MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
// Client-side cancellation or end
if let Some(tx) = streams.remove(&frame.stream_id) {
let _ = dispatch_stream_frame(&tx, frame).await;
if let Some(target) = streams.remove(&frame.stream_id) {
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
if draining && streams.is_empty() {
info!("tunnel drained after stream termination");
break None;
@@ -320,6 +359,35 @@ where
break None;
}
MsgType::WindowUpdate => {
if let Ok(payload) = serde_json::from_slice::<
aether_contracts::tunnel::WindowUpdatePayload,
>(&frame.payload)
{
if let Some(target) = streams.get(&frame.stream_id) {
target.response_window.add_credit(payload.delta_bytes);
}
}
debug!(
msg_type = ?frame.msg_type,
stream_id = frame.stream_id,
"received tunnel protocol v3 WINDOW_UPDATE frame"
);
}
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
debug!(
msg_type = ?frame.msg_type,
stream_id = frame.stream_id,
"received tunnel protocol v3 control frame"
);
}
MsgType::ConnectionClose => {
info!("received CONNECTION_CLOSE");
break None;
}
_ => {
debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type");
}
@@ -329,10 +397,6 @@ 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 {
let closed_streams = prune_closed_stream_senders(&mut streams);
if closed_streams > 0 {
debug!(closed_streams, "removed closed request body stream senders");
}
handler_handles.retain(|h| !h.is_finished());
frames_since_cleanup = 0;
if draining && streams.is_empty() {
@@ -408,9 +472,10 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
}
}
fn prune_closed_stream_senders(streams: &mut HashMap<u32, mpsc::Sender<Frame>>) -> usize {
#[cfg(test)]
fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>) -> usize {
let before = streams.len();
streams.retain(|_, tx| !tx.is_closed());
streams.retain(|_, target| !target.body_tx.is_closed());
before.saturating_sub(streams.len())
}
@@ -493,7 +558,22 @@ mod tests {
let (closed_tx, closed_rx) = mpsc::channel::<Frame>(1);
let (open_tx, _open_rx) = mpsc::channel::<Frame>(1);
drop(closed_rx);
let mut streams = HashMap::from([(7, closed_tx), (9, open_tx)]);
let mut streams = HashMap::from([
(
7,
StreamDispatchTarget {
body_tx: closed_tx,
response_window: Arc::new(StreamSendWindow::new(1024)),
},
),
(
9,
StreamDispatchTarget {
body_tx: open_tx,
response_window: Arc::new(StreamSendWindow::new(1024)),
},
),
]);
let removed = prune_closed_stream_senders(&mut streams);
+4
View File
@@ -548,6 +548,10 @@ mod tests {
tunnel_reconnect_max_ms: 250,
tunnel_ping_interval_ms: 1_000,
tunnel_max_streams: Some(8),
tunnel_profile: crate::config::TunnelProfileArg::Lite,
tunnel_stream_initial_window_bytes:
crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES,
tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS,
tunnel_connect_timeout_ms: 2_000,
tunnel_ipv4_only: false,
tunnel_ipv6_only: false,
+233 -17
View File
@@ -25,7 +25,7 @@ use crate::upstream_client;
use super::protocol::{
compress_payload, decompress_if_gzip, flags, raw_payload, Frame as TunnelFrame, MsgType,
RequestMeta, ResponseMeta,
RequestMeta, ResetStreamPayload, ResponseMeta,
};
use super::writer::FrameSender;
@@ -35,10 +35,101 @@ const MAX_CHUNK_SIZE: usize = 32 * 1024;
/// Timeout for sending a single frame to the writer channel.
/// Control frames are allowed a short wait; body frames fail fast.
const CONTROL_FRAME_SEND_TIMEOUT: Duration = Duration::from_millis(250);
const FLOW_CONTROL_WAIT_TIMEOUT: Duration = Duration::from_secs(5);
const SLOW_STREAM_LOG_THRESHOLD: Duration = Duration::from_secs(2);
const SUCCESS_LOG_SAMPLE_MODULO: u32 = 256;
const REQUEST_BODY_SPOOL_QUEUE_CAPACITY: usize = 64;
#[derive(Debug)]
pub(crate) struct StreamSendWindow {
available: Mutex<u64>,
notify: Notify,
}
impl StreamSendWindow {
pub(crate) fn new(initial_window_bytes: u32) -> Self {
Self {
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
notify: Notify::new(),
}
}
pub(crate) fn add_credit(&self, delta_bytes: u32) {
if delta_bytes == 0 {
return;
}
let mut available = self.available.lock().expect("stream window lock poisoned");
*available = available.saturating_add(u64::from(delta_bytes));
drop(available);
self.notify.notify_waiters();
}
async fn acquire(&self, bytes: usize, timeout: Duration) -> Result<Duration, ()> {
if bytes == 0 {
return Ok(Duration::ZERO);
}
let requested = bytes as u64;
let started_at = Instant::now();
loop {
{
let mut available = self.available.lock().expect("stream window lock poisoned");
if *available >= requested {
*available -= requested;
return Ok(started_at.elapsed());
}
}
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(());
};
if tokio::time::timeout(remaining, self.notify.notified())
.await
.is_err()
{
return Err(());
}
}
}
}
fn stream_reset_message(frame: &TunnelFrame) -> String {
if frame.msg_type == MsgType::ResetStream {
if let Ok(payload) = serde_json::from_slice::<ResetStreamPayload>(&frame.payload) {
return payload.reason;
}
}
String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string())
}
fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) {
if bytes == 0 {
return;
}
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"
);
}
}
/// Minimum allowed upstream request timeout (milliseconds).
const MIN_TIMEOUT_MS: u64 = 1;
/// Maximum allowed upstream request timeout (milliseconds).
@@ -491,10 +582,12 @@ fn buffered_request_body(body: Bytes) -> upstream_client::UpstreamRequestBody {
// longer coupled to upstream body polling. Redirect replay still reuses a full
// in-memory copy when the request body completes within budget.
fn prepare_request_body(
stream_id: u32,
body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>,
deadline: Instant,
replay_budget_bytes: usize,
frame_tx: FrameSender,
) -> PreparedRequestBody {
let (spool_tx, spool_rx) = mpsc::channel(REQUEST_BODY_SPOOL_QUEUE_CAPACITY);
let replay_state = if replay_budget_bytes == 0 {
@@ -508,11 +601,13 @@ fn prepare_request_body(
};
tokio::spawn(spool_request_body(
stream_id,
body_rx,
spool_tx,
replay_state,
body_size,
deadline,
frame_tx,
));
PreparedRequestBody {
@@ -537,10 +632,12 @@ fn prepare_bodyless_request_body(
}
async fn collect_request_body_for_replay(
stream_id: u32,
mut body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>,
deadline: Instant,
replay_budget_bytes: usize,
frame_tx: &FrameSender,
) -> Result<Bytes, String> {
let mut body = BytesMut::new();
@@ -565,6 +662,7 @@ async fn collect_request_body_for_replay(
));
}
body_size.fetch_add(payload.len(), Ordering::Relaxed);
try_send_window_update(frame_tx, stream_id, payload.len());
body.extend_from_slice(&payload);
}
@@ -572,9 +670,8 @@ async fn collect_request_body_for_replay(
return Ok(body.freeze());
}
}
MsgType::StreamError => {
return Err(String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string()));
MsgType::StreamError | MsgType::ResetStream => {
return Err(stream_reset_message(&frame));
}
MsgType::StreamEnd => return Ok(body.freeze()),
_ => continue,
@@ -648,11 +745,13 @@ fn timeout_duration_from_legacy_secs(secs: u64) -> Duration {
}
async fn spool_request_body(
stream_id: u32,
mut body_rx: mpsc::Receiver<TunnelFrame>,
mut spool_tx: mpsc::Sender<SpoolBodyEvent>,
replay_state: Option<Arc<RequestBodyReplayState>>,
body_size: Arc<AtomicUsize>,
deadline: Instant,
frame_tx: FrameSender,
) {
loop {
let frame = match recv_body_frame_with_deadline(&mut body_rx, deadline).await {
@@ -692,6 +791,7 @@ 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());
}
@@ -714,9 +814,8 @@ async fn spool_request_body(
return;
}
}
MsgType::StreamError => {
let message = String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string());
MsgType::StreamError | MsgType::ResetStream => {
let message = stream_reset_message(&frame);
if let Some(state) = &replay_state {
state.fail(message.clone());
}
@@ -932,6 +1031,40 @@ async fn execute_upstream_request(
})
}
async fn acquire_response_credit(
response_window: &StreamSendWindow,
frame_tx: &FrameSender,
stream_id: u32,
bytes: usize,
) -> bool {
match response_window
.acquire(bytes, FLOW_CONTROL_WAIT_TIMEOUT)
.await
{
Ok(waited) => {
if waited > Duration::from_millis(1) {
debug!(
stream_id,
bytes,
waited_ms = waited.as_millis() as u64,
"waited for tunnel response flow-control credit"
);
}
true
}
Err(()) => {
warn!(
stream_id,
bytes,
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
"response flow-control window timeout"
);
send_reset_stream(frame_tx, stream_id, "response_flow_control_timeout").await;
false
}
}
}
#[allow(clippy::too_many_arguments)]
async fn relay_upstream_response<B>(
server: &ServerContext,
@@ -939,6 +1072,7 @@ async fn relay_upstream_response<B>(
method: &hyper::Method,
request_url: &url::Url,
frame_tx: &FrameSender,
response_window: &StreamSendWindow,
response: hyper::Response<B>,
total_dns_ms: u64,
total_elapsed: Duration,
@@ -1068,6 +1202,11 @@ where
Ok(chunk) => {
if chunk.len() <= MAX_CHUNK_SIZE {
let (payload, extra_flags) = raw_payload(chunk);
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
.await
{
return Some(total_elapsed);
}
if !send_frame(
frame_tx,
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
@@ -1094,6 +1233,16 @@ where
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
let slice = chunk.slice(offset..end);
let (payload, extra_flags) = raw_payload(slice);
if !acquire_response_credit(
response_window,
frame_tx,
stream_id,
payload.len(),
)
.await
{
return Some(total_elapsed);
}
if !send_frame(
frame_tx,
TunnelFrame::new(
@@ -1213,6 +1362,7 @@ pub async fn handle_stream(
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: FrameSender,
response_window: Arc<StreamSendWindow>,
) {
let request_method = parse_request_method(&meta.method);
let request_url = url::Url::parse(&meta.url).ok();
@@ -1244,8 +1394,17 @@ pub async fn handle_stream(
server.active_connections.fetch_add(1, Ordering::Release);
let connect_elapsed =
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx, permit).await;
let connect_elapsed = handle_stream_inner(
&state,
&server,
stream_id,
meta,
body_rx,
&frame_tx,
response_window.as_ref(),
permit,
)
.await;
server.active_connections.fetch_sub(1, Ordering::Release);
if let Some(d) = connect_elapsed {
@@ -1264,18 +1423,21 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
);
if is_body_frame {
match tx.try_send(frame) {
Ok(()) => true,
Err(QueueSendError::Full(_)) => {
match tokio::time::timeout(FLOW_CONTROL_WAIT_TIMEOUT, tx.send(frame)).await {
Ok(Ok(())) => true,
Ok(Err(QueueSendError::Closed(_))) | Err(_) => {
warn!(
stream_id,
msg_type = ?msg_type,
flags = flags,
"writer channel full for body frame, abandoning stream"
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
"writer channel stalled for body frame, abandoning stream"
);
false
}
Err(QueueSendError::Closed(_)) => false,
Ok(Err(QueueSendError::Full(_))) => {
unreachable!("bounded queue send should not report full")
}
}
} else {
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
@@ -1304,6 +1466,7 @@ async fn handle_stream_inner(
meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: &FrameSender,
response_window: &StreamSendWindow,
mut admission_permit: Option<AdmissionPermit>,
) -> Option<Duration> {
let mut current_method: hyper::Method = parse_request_method(&meta.method);
@@ -1357,10 +1520,12 @@ async fn handle_stream_inner(
};
let mut prepared_body = if can_buffer_redirect_body {
let buffered_body = match collect_request_body_for_replay(
stream_id,
body_rx,
Arc::clone(&request_body_size),
first_byte_deadline,
replay_budget_bytes,
frame_tx,
)
.await
{
@@ -1388,10 +1553,12 @@ async fn handle_stream_inner(
}
} else if request_has_body {
prepare_request_body(
stream_id,
body_rx,
Arc::clone(&request_body_size),
first_byte_deadline,
0,
frame_tx.clone(),
)
} else {
prepare_bodyless_request_body(body_rx, follow_redirects)
@@ -1472,6 +1639,7 @@ async fn handle_stream_inner(
&current_method,
&current_url,
frame_tx,
response_window,
response_ctx.response,
total_dns_ms,
overall_start.elapsed(),
@@ -1512,6 +1680,7 @@ async fn handle_stream_inner(
&current_method,
&current_url,
frame_tx,
response_window,
response_ctx.response,
total_dns_ms,
overall_start.elapsed(),
@@ -1568,6 +1737,7 @@ async fn handle_stream_inner(
&current_method,
&current_url,
frame_tx,
response_window,
response_ctx.response,
total_dns_ms,
overall_start.elapsed(),
@@ -1596,6 +1766,18 @@ async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
.await;
}
async fn send_reset_stream(tx: &FrameSender, stream_id: u32, reason: &str) {
let payload = serde_json::to_vec(&ResetStreamPayload {
reason: reason.to_string(),
})
.expect("reset stream payload should serialize");
let _ = send_frame(
tx,
TunnelFrame::new(stream_id, MsgType::ResetStream, 0, Bytes::from(payload)),
)
.await;
}
#[cfg(test)]
fn build_streaming_request_body(
body_rx: mpsc::Receiver<TunnelFrame>,
@@ -1676,9 +1858,8 @@ fn build_prefixed_request_body(
(body_rx, body_size, end_stream),
));
}
MsgType::StreamError => {
let message = String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string());
MsgType::StreamError | MsgType::ResetStream => {
let message = stream_reset_message(&frame);
return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
}
MsgType::StreamEnd => return None,
@@ -1820,12 +2001,15 @@ mod tests {
#[tokio::test]
async fn prepare_request_body_streams_immediately_and_replays_after_completion() {
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(
1,
rx,
Arc::clone(&body_size),
Instant::now() + Duration::from_secs(1),
1024,
frame_tx.clone(),
);
let mut body = prepared
.first_request_body
@@ -1888,6 +2072,19 @@ mod tests {
);
assert!(replay.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 11);
let window_update_bytes = collect_emitted_frames(frame_tx, sent, writer_handle)
.await
.into_iter()
.filter(|frame| frame.msg_type == MsgType::WindowUpdate)
.filter_map(|frame| {
serde_json::from_slice::<aether_contracts::tunnel::WindowUpdatePayload>(
&frame.payload,
)
.ok()
})
.map(|payload| payload.delta_bytes as usize)
.sum::<usize>();
assert_eq!(window_update_bytes, 11);
}
#[test]
@@ -2084,6 +2281,7 @@ mod tests {
meta,
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2145,6 +2343,7 @@ mod tests {
meta,
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2179,6 +2378,7 @@ mod tests {
.status(StatusCode::OK)
.body(body)
.expect("response");
let response_window = test_response_window();
relay_upstream_response(
&server,
@@ -2186,6 +2386,7 @@ mod tests {
&hyper::Method::GET,
&request_url,
&frame_tx,
response_window.as_ref(),
response,
0,
Duration::ZERO,
@@ -2222,6 +2423,7 @@ mod tests {
.status(StatusCode::OK)
.body(body)
.expect("response");
let response_window = test_response_window();
relay_upstream_response(
&server,
@@ -2229,6 +2431,7 @@ mod tests {
&hyper::Method::GET,
&request_url,
&frame_tx,
response_window.as_ref(),
response,
0,
Duration::ZERO,
@@ -2329,6 +2532,7 @@ mod tests {
meta,
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2393,6 +2597,7 @@ mod tests {
meta,
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2477,6 +2682,7 @@ mod tests {
meta,
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2515,6 +2721,7 @@ mod tests {
sample_request_meta(),
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
@@ -2561,6 +2768,7 @@ mod tests {
sample_request_meta(),
body_rx,
frame_tx.clone(),
test_response_window(),
)
.await;
@@ -2729,6 +2937,10 @@ mod tests {
tunnel_reconnect_max_ms: 30_000,
tunnel_ping_interval_ms: 15_000,
tunnel_max_streams: Some(8),
tunnel_profile: crate::config::TunnelProfileArg::Lite,
tunnel_stream_initial_window_bytes:
crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES,
tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS,
tunnel_connect_timeout_ms: 15_000,
tunnel_ipv4_only: false,
tunnel_ipv6_only: false,
@@ -2793,6 +3005,10 @@ mod tests {
(frame_tx, sent, handle)
}
fn test_response_window() -> Arc<StreamSendWindow> {
Arc::new(StreamSendWindow::new(u32::MAX))
}
struct StreamResult {
response: Option<ResponseMeta>,
body: Bytes,
+7 -1
View File
@@ -188,7 +188,13 @@ fn classify_frame_priority(frame: &Frame) -> FramePriority {
| MsgType::Pong
| MsgType::GoAway
| MsgType::HeartbeatData
| MsgType::HeartbeatAck => FramePriority::High,
| MsgType::HeartbeatAck
| MsgType::Hello
| MsgType::Settings
| MsgType::WindowUpdate
| MsgType::ResetStream
| MsgType::ConnectionClose
| MsgType::LoadReport => FramePriority::High,
MsgType::RequestHeaders
| MsgType::RequestBody
| MsgType::ResponseBody