mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 19:29:50 +08:00
Refactor tunnel stability protocol
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
¤t_method,
|
||||
¤t_url,
|
||||
frame_tx,
|
||||
response_window,
|
||||
response_ctx.response,
|
||||
total_dns_ms,
|
||||
overall_start.elapsed(),
|
||||
@@ -1512,6 +1680,7 @@ async fn handle_stream_inner(
|
||||
¤t_method,
|
||||
¤t_url,
|
||||
frame_tx,
|
||||
response_window,
|
||||
response_ctx.response,
|
||||
total_dns_ms,
|
||||
overall_start.elapsed(),
|
||||
@@ -1568,6 +1737,7 @@ async fn handle_stream_inner(
|
||||
¤t_method,
|
||||
¤t_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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user