mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 00:17:45 +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:
@@ -47,3 +47,4 @@ uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
aether-gateway = { workspace = true, features = ["testkit"] }
|
||||
tokio = { version = "1", features = ["test-util"] }
|
||||
|
||||
@@ -4,6 +4,15 @@ Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道
|
||||
|
||||
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
||||
|
||||
## 流式传输与升级注意事项
|
||||
|
||||
- 协议 v3 连接在 `HELLO` / `SETTINGS` 协商后才接收业务请求。实际双向流窗口取 gateway 与 agent 配置的较小值,信用更新阈值不超过该窗口的四分之一;单帧也不会超过协商窗口。
|
||||
- 响应缓冲按字节限额并合并小帧,结束和错误状态独立保存。慢消费者不会阻塞同一隧道其他流的读取;超出窗口或缓冲预算的流会被明确终止,不会静默截断。
|
||||
- 信用更新在消费数据后可靠入队;启用重定向重放时,进入有界重放缓存也视为请求体消费。持续无法投递关键控制帧时会关闭连接并向在途请求报告错误。
|
||||
- 客户端取消会终止对应上游请求,断连会回收 session 的 writer、heartbeat 和请求任务。正常 drain 在配置期限内继续处理已有流,期限到达后终止残留任务。
|
||||
- 建议先升级 gateway,再升级 agent。既有 v3 agent 已发送 `HELLO` / `SETTINGS`,可连接新 gateway;自定义 v3 节点必须完成这两步握手。协议 v1/v2 保留旧握手。与旧 gateway 混用时应保持默认窗口配置,不能依赖旧 gateway 应用新的窗口协商。
|
||||
- 自动重连恢复后续请求,不会自动续传已经输出的 SSE,也不会无条件重放已经发送的请求。
|
||||
|
||||
## 安装
|
||||
|
||||
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
||||
|
||||
@@ -764,6 +764,13 @@ impl Config {
|
||||
if self.tunnel_stream_initial_window_bytes == 0 {
|
||||
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
||||
}
|
||||
if u64::from(self.tunnel_stream_initial_window_bytes)
|
||||
> aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||
{
|
||||
anyhow::bail!(
|
||||
"tunnel_stream_initial_window_bytes exceeds the maximum tunnel payload size"
|
||||
);
|
||||
}
|
||||
if self.tunnel_drain_deadline_ms == 0 {
|
||||
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
||||
}
|
||||
|
||||
@@ -203,19 +203,34 @@ pub async fn connect_and_run(
|
||||
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
||||
|
||||
// Spawn writer task (with WebSocket ping keepalive)
|
||||
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||
let (frame_tx, writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||
ws_sink,
|
||||
ping_interval,
|
||||
Some(Arc::clone(&server.tunnel_metrics)),
|
||||
security.clone(),
|
||||
);
|
||||
let mut writer_handle = super::task::SessionTask::new(writer_handle);
|
||||
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
||||
let drain_signal = spawn_drain_signal(
|
||||
let (session_drain_tx, session_drain_rx) = watch::channel(*drain.borrow());
|
||||
let forward_drain_tx = session_drain_tx.clone();
|
||||
let mut external_drain = drain;
|
||||
let forward_drain = super::task::SessionTask::new(tokio::spawn(async move {
|
||||
loop {
|
||||
if *external_drain.borrow() {
|
||||
let _ = forward_drain_tx.send(true);
|
||||
break;
|
||||
}
|
||||
if external_drain.changed().await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}));
|
||||
let drain_signal = super::task::SessionTask::new(spawn_drain_signal(
|
||||
conn_idx,
|
||||
frame_tx.clone(),
|
||||
drain.clone(),
|
||||
session_drain_rx.clone(),
|
||||
state.config.tunnel_drain_deadline_ms,
|
||||
);
|
||||
));
|
||||
|
||||
// Spawn heartbeat task (only for primary connection to avoid
|
||||
// resetting shared atomic metrics via swap(0))
|
||||
@@ -237,16 +252,19 @@ pub async fn connect_and_run(
|
||||
// ensures we detect this and trigger a reconnect promptly.
|
||||
let state_clone = Arc::clone(state);
|
||||
let server_clone = Arc::clone(server);
|
||||
let outcome = tokio::select! {
|
||||
result = dispatcher::run_with_security(
|
||||
let outcome = {
|
||||
let dispatch = dispatcher::run_with_security(
|
||||
state_clone,
|
||||
server_clone,
|
||||
ws_read,
|
||||
frame_tx.clone(),
|
||||
hb_handle,
|
||||
drain.clone(),
|
||||
session_drain_rx,
|
||||
security.clone(),
|
||||
) => {
|
||||
);
|
||||
tokio::pin!(dispatch);
|
||||
tokio::select! {
|
||||
result = &mut dispatch => {
|
||||
match result {
|
||||
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
||||
Err(e) => {
|
||||
@@ -258,6 +276,8 @@ pub async fn connect_and_run(
|
||||
}
|
||||
}
|
||||
writer_result = &mut writer_handle => {
|
||||
frame_tx.close();
|
||||
let _ = tokio::time::timeout(Duration::from_secs(1), &mut dispatch).await;
|
||||
match writer_result {
|
||||
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
||||
Err(e) => {
|
||||
@@ -278,24 +298,37 @@ pub async fn connect_and_run(
|
||||
}
|
||||
_ = shutdown.changed() => {
|
||||
debug!("shutdown during tunnel dispatch");
|
||||
let _ = session_drain_tx.send(true);
|
||||
let deadline = Duration::from_millis(state.config.tunnel_drain_deadline_ms).saturating_add(Duration::from_secs(1));
|
||||
let _ = tokio::time::timeout(deadline, &mut dispatch).await;
|
||||
Ok(TunnelOutcome::Shutdown)
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Drop our sender; the writer will exit once all stream handler clones
|
||||
// are also dropped (i.e. after they finish their in-flight work).
|
||||
drop(frame_tx);
|
||||
forward_drain.abort();
|
||||
let _ = forward_drain.await;
|
||||
if !drain_signal.is_finished() {
|
||||
drain_signal.abort();
|
||||
let _ = drain_signal.await;
|
||||
}
|
||||
|
||||
// Wait for the writer task to finish with a generous timeout — the
|
||||
// dispatcher already waits up to 30s for stream handlers, so 35s here
|
||||
// covers that plus a small margin.
|
||||
// Skip if the writer already exited (the select branch that fired).
|
||||
if !writer_handle.is_finished() {
|
||||
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
|
||||
let flush_timeout = if *session_drain_tx.borrow() {
|
||||
Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||
} else {
|
||||
Duration::from_secs(1)
|
||||
};
|
||||
if tokio::time::timeout(flush_timeout, &mut writer_handle)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
writer_handle.abort();
|
||||
let _ = writer_handle.await;
|
||||
}
|
||||
}
|
||||
|
||||
let connected_for = connected_at.elapsed();
|
||||
|
||||
@@ -9,7 +9,7 @@ use std::time::Duration;
|
||||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio::task::{AbortHandle, JoinSet};
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
@@ -41,13 +41,25 @@ impl AsRef<[u8]> for BudgetedFramePayload {
|
||||
enum StreamDispatchStatus {
|
||||
Delivered,
|
||||
Closed,
|
||||
TimedOut,
|
||||
Congested,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct StreamDispatchTarget {
|
||||
body_tx: mpsc::Sender<Frame>,
|
||||
response_window: Arc<StreamSendWindow>,
|
||||
handler: Option<AbortHandle>,
|
||||
}
|
||||
|
||||
struct StreamCompletion {
|
||||
stream_id: u32,
|
||||
finished_tx: mpsc::UnboundedSender<u32>,
|
||||
}
|
||||
|
||||
impl Drop for StreamCompletion {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.finished_tx.send(self.stream_id);
|
||||
}
|
||||
}
|
||||
|
||||
/// A request stream is identified by a non-zero id and may only be opened
|
||||
@@ -109,7 +121,7 @@ where
|
||||
// reopen the same id and bypass the stream admission limit.
|
||||
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
||||
// Track spawned stream handlers so we can wait for them on shutdown
|
||||
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
|
||||
let mut handler_handles = JoinSet::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;
|
||||
@@ -121,30 +133,48 @@ where
|
||||
// Track last time we received any data to detect stale connections
|
||||
let mut last_data_at = tokio::time::Instant::now();
|
||||
let mut draining = *drain.borrow();
|
||||
let mut drain_open = true;
|
||||
let mut drain_deadline = draining.then(|| {
|
||||
tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||
});
|
||||
let mut initial_window_bytes = state.config.tunnel_stream_initial_window_bytes;
|
||||
let mut close_rx = frame_tx.subscribe_close();
|
||||
|
||||
let read_err = loop {
|
||||
if *close_rx.borrow() {
|
||||
break None;
|
||||
}
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after in-flight streams completed");
|
||||
break None;
|
||||
}
|
||||
|
||||
let msg_result = tokio::select! {
|
||||
_ = close_rx.changed() => break None,
|
||||
msg = ws_stream.next() => {
|
||||
match msg {
|
||||
Some(r) => r,
|
||||
None => break None,
|
||||
}
|
||||
}
|
||||
changed = drain.changed() => {
|
||||
changed = drain.changed(), if drain_open => {
|
||||
if changed.is_err() {
|
||||
drain_open = false;
|
||||
continue;
|
||||
}
|
||||
if *drain.borrow() {
|
||||
info!("tunnel drain requested, waiting for in-flight streams");
|
||||
draining = true;
|
||||
drain_deadline.get_or_insert_with(|| tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms));
|
||||
}
|
||||
continue;
|
||||
}
|
||||
_ = async {
|
||||
match drain_deadline {
|
||||
Some(deadline) => tokio::time::sleep_until(deadline).await,
|
||||
None => std::future::pending().await,
|
||||
}
|
||||
} => break None,
|
||||
finished = handler_finished_rx.recv() => {
|
||||
if let Some(stream_id) = finished {
|
||||
active_handler_ids.remove(&stream_id);
|
||||
@@ -238,20 +268,7 @@ where
|
||||
continue;
|
||||
}
|
||||
if draining {
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from("tunnel draining"),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped during drain"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(&frame_tx, frame.stream_id, "tunnel draining");
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -263,6 +280,11 @@ where
|
||||
Ok(p) => p,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"invalid request metadata",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
@@ -270,21 +292,11 @@ where
|
||||
Ok(m) => m,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
||||
// Use try_send to avoid blocking the read loop
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from(format!("invalid request metadata: {e}")),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"invalid request metadata",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
@@ -294,33 +306,27 @@ where
|
||||
stream_id = frame.stream_id,
|
||||
"max concurrent streams reached"
|
||||
);
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from("max concurrent streams reached"),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"max concurrent streams reached",
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Create body channel and spawn handler
|
||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
|
||||
let response_window = Arc::new(StreamSendWindow::new(
|
||||
state.config.tunnel_stream_initial_window_bytes,
|
||||
));
|
||||
let body_capacity = (initial_window_bytes as usize)
|
||||
.div_ceil(32 * 1024)
|
||||
.saturating_add(1)
|
||||
.max(64);
|
||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(body_capacity);
|
||||
let response_window = Arc::new(StreamSendWindow::new(initial_window_bytes));
|
||||
streams.insert(
|
||||
frame.stream_id,
|
||||
StreamDispatchTarget {
|
||||
body_tx,
|
||||
response_window: Arc::clone(&response_window),
|
||||
handler: None,
|
||||
},
|
||||
);
|
||||
active_handler_ids.insert(frame.stream_id);
|
||||
@@ -329,9 +335,13 @@ where
|
||||
let state_clone = Arc::clone(&state);
|
||||
let server_clone = Arc::clone(&server);
|
||||
let tx_clone = frame_tx.clone();
|
||||
let finished_tx = handler_finished_tx.clone();
|
||||
let sid = frame.stream_id;
|
||||
let handle = tokio::spawn(async move {
|
||||
let completion = StreamCompletion {
|
||||
stream_id: sid,
|
||||
finished_tx: handler_finished_tx.clone(),
|
||||
};
|
||||
let handle = handler_handles.spawn(async move {
|
||||
let _completion = completion;
|
||||
stream_handler::handle_stream(
|
||||
state_clone,
|
||||
server_clone,
|
||||
@@ -342,9 +352,8 @@ where
|
||||
response_window,
|
||||
)
|
||||
.await;
|
||||
let _ = finished_tx.send(sid);
|
||||
});
|
||||
handler_handles.push(handle);
|
||||
streams.get_mut(&sid).expect("new stream exists").handler = Some(handle);
|
||||
|
||||
if request_headers_end_stream {
|
||||
if let Some(target) = streams.get(&sid) {
|
||||
@@ -365,19 +374,21 @@ where
|
||||
let is_end = frame.is_end_stream();
|
||||
let sid = frame.stream_id;
|
||||
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||
if dispatch != StreamDispatchStatus::Delivered {
|
||||
streams.remove(&sid);
|
||||
if dispatch == StreamDispatchStatus::TimedOut {
|
||||
server.tunnel_metrics.record_error(
|
||||
"stream_dispatch_timeout",
|
||||
&format!("request body dispatch timed out for stream {}", sid),
|
||||
);
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
sid,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
if dispatch == StreamDispatchStatus::Congested {
|
||||
if let Some(target) = streams.remove(&sid) {
|
||||
if let Some(handler) = target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
}
|
||||
server.tunnel_metrics.record_error(
|
||||
"stream_dispatch_timeout",
|
||||
&format!("request body dispatch congested for stream {}", sid),
|
||||
);
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
sid,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
||||
{
|
||||
info!("tunnel drained after request body completion");
|
||||
@@ -387,10 +398,29 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
|
||||
MsgType::StreamEnd => {
|
||||
if let Some(target) = streams.get(&frame.stream_id) {
|
||||
if dispatch_stream_frame(&target.body_tx, frame.clone()).await
|
||||
== StreamDispatchStatus::Congested
|
||||
{
|
||||
if let Some(handler) = &target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
try_send_stream_error(
|
||||
&frame_tx,
|
||||
frame.stream_id,
|
||||
"tunnel request body dispatch stalled",
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::StreamError | MsgType::ResetStream => {
|
||||
// Client-side cancellation or end
|
||||
if let Some(target) = streams.remove(&frame.stream_id) {
|
||||
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||
if let Some(handler) = target.handler {
|
||||
handler.abort();
|
||||
}
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after stream termination");
|
||||
break None;
|
||||
@@ -409,12 +439,22 @@ where
|
||||
}
|
||||
|
||||
MsgType::HeartbeatAck => {
|
||||
heartbeat.on_ack(frame.payload).await;
|
||||
heartbeat.on_ack(frame.payload);
|
||||
}
|
||||
|
||||
MsgType::GoAway => {
|
||||
info!("received GOAWAY");
|
||||
break None;
|
||||
draining = true;
|
||||
let deadline_ms =
|
||||
serde_json::from_slice::<aether_contracts::tunnel::GoAwayPayload>(
|
||||
&frame.payload,
|
||||
)
|
||||
.map(|payload| payload.drain_deadline_ms)
|
||||
.unwrap_or(state.config.tunnel_drain_deadline_ms)
|
||||
.min(state.config.tunnel_drain_deadline_ms);
|
||||
drain_deadline.get_or_insert_with(|| {
|
||||
tokio::time::Instant::now() + Duration::from_millis(deadline_ms)
|
||||
});
|
||||
}
|
||||
|
||||
MsgType::WindowUpdate => {
|
||||
@@ -433,7 +473,31 @@ where
|
||||
);
|
||||
}
|
||||
|
||||
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
|
||||
MsgType::Settings => {
|
||||
if frame.stream_id != 0 || frame.flags != 0 {
|
||||
break None;
|
||||
}
|
||||
let settings = serde_json::from_slice::<aether_contracts::tunnel::SettingsPayload>(
|
||||
&frame.payload,
|
||||
)
|
||||
.ok()
|
||||
.filter(|settings| settings.is_valid());
|
||||
let Some(settings) = settings else {
|
||||
warn!("invalid tunnel SETTINGS");
|
||||
break None;
|
||||
};
|
||||
if !streams.is_empty()
|
||||
&& settings.initial_stream_window_bytes != initial_window_bytes
|
||||
{
|
||||
warn!("tunnel SETTINGS changed with active streams");
|
||||
break None;
|
||||
}
|
||||
initial_window_bytes = settings
|
||||
.initial_stream_window_bytes
|
||||
.min(state.config.tunnel_stream_initial_window_bytes);
|
||||
}
|
||||
|
||||
MsgType::Hello | MsgType::LoadReport => {
|
||||
debug!(
|
||||
msg_type = ?frame.msg_type,
|
||||
stream_id = frame.stream_id,
|
||||
@@ -455,7 +519,7 @@ where
|
||||
// Trigger every 64 frames OR when the count exceeds max_streams.
|
||||
frames_since_cleanup += 1;
|
||||
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
||||
handler_handles.retain(|h| !h.is_finished());
|
||||
while handler_handles.try_join_next().is_some() {}
|
||||
frames_since_cleanup = 0;
|
||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||
info!("tunnel drained after cleanup");
|
||||
@@ -467,9 +531,7 @@ where
|
||||
// Drop body senders so stream handlers waiting on body_rx will unblock
|
||||
streams.clear();
|
||||
|
||||
// Wait for active stream handlers to finish so their frame_tx clones
|
||||
// are dropped before the writer closes the sink.
|
||||
drain_handlers(handler_handles).await;
|
||||
handler_handles.shutdown().await;
|
||||
|
||||
match read_err {
|
||||
Some(e) => Err(e.into()),
|
||||
@@ -478,30 +540,13 @@ where
|
||||
}
|
||||
|
||||
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
||||
let stream_id = frame.stream_id;
|
||||
let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async {
|
||||
let frame = attach_request_body_queue_budget(frame).await?;
|
||||
tx.send(frame).await.ok()?;
|
||||
Some(())
|
||||
})
|
||||
.await;
|
||||
match dispatched {
|
||||
Ok(Some(())) => StreamDispatchStatus::Delivered,
|
||||
Ok(None) => {
|
||||
warn!(
|
||||
stream_id,
|
||||
"stream handler channel or request body budget closed while dispatching tunnel frame"
|
||||
);
|
||||
StreamDispatchStatus::Closed
|
||||
}
|
||||
Err(_) => {
|
||||
warn!(
|
||||
stream_id,
|
||||
timeout_ms = stream_frame_dispatch_timeout().as_millis(),
|
||||
"stream handler channel blocked while dispatching tunnel frame"
|
||||
);
|
||||
StreamDispatchStatus::TimedOut
|
||||
}
|
||||
let Some(frame) = attach_request_body_queue_budget(frame).await else {
|
||||
return StreamDispatchStatus::Congested;
|
||||
};
|
||||
match tx.try_send(frame) {
|
||||
Ok(()) => StreamDispatchStatus::Delivered,
|
||||
Err(mpsc::error::TrySendError::Closed(_)) => StreamDispatchStatus::Closed,
|
||||
Err(mpsc::error::TrySendError::Full(_)) => StreamDispatchStatus::Congested,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -523,7 +568,7 @@ async fn attach_request_body_queue_budget_with(
|
||||
return Some(frame);
|
||||
}
|
||||
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
||||
let permit = budget.acquire_many_owned(permits).await.ok()?;
|
||||
let permit = budget.try_acquire_many_owned(permits).ok()?;
|
||||
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
||||
bytes: frame.payload,
|
||||
_permit: permit,
|
||||
@@ -549,20 +594,6 @@ fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option<u32>
|
||||
u32::try_from(retained_bytes).ok()
|
||||
}
|
||||
|
||||
/// Bound how long a single stream handler is allowed to block the shared
|
||||
/// WebSocket read loop while receiving request-body frames.
|
||||
fn stream_frame_dispatch_timeout() -> Duration {
|
||||
#[cfg(test)]
|
||||
{
|
||||
Duration::from_millis(25)
|
||||
}
|
||||
|
||||
#[cfg(not(test))]
|
||||
{
|
||||
Duration::from_millis(500)
|
||||
}
|
||||
}
|
||||
|
||||
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
@@ -573,6 +604,7 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
frame_tx.close();
|
||||
warn!(
|
||||
stream_id,
|
||||
"writer channel full, StreamError dropped while aborting stalled stream"
|
||||
@@ -587,21 +619,6 @@ fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>)
|
||||
before.saturating_sub(streams.len())
|
||||
}
|
||||
|
||||
/// Wait for all active stream handlers to finish (with a timeout).
|
||||
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
|
||||
if handles.is_empty() {
|
||||
return;
|
||||
}
|
||||
let count = handles.len();
|
||||
debug!(count, "waiting for active stream handlers to finish");
|
||||
let _ = tokio::time::timeout(Duration::from_secs(30), async {
|
||||
for h in handles {
|
||||
let _ = h.await;
|
||||
}
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -633,7 +650,7 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
stalled_send.await.expect("dispatch task should join"),
|
||||
StreamDispatchStatus::TimedOut
|
||||
StreamDispatchStatus::Congested
|
||||
);
|
||||
|
||||
let retained = rx
|
||||
@@ -737,6 +754,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: closed_tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
),
|
||||
(
|
||||
@@ -744,6 +762,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: open_tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
),
|
||||
]);
|
||||
@@ -763,6 +782,7 @@ mod tests {
|
||||
StreamDispatchTarget {
|
||||
body_tx: tx,
|
||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||
handler: None,
|
||||
},
|
||||
)]);
|
||||
let mut active_handler_ids = HashSet::from([7]);
|
||||
|
||||
@@ -31,14 +31,22 @@ enum AckDecision {
|
||||
}
|
||||
|
||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||
#[derive(Clone)]
|
||||
pub struct HeartbeatHandle {
|
||||
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
||||
task: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl HeartbeatHandle {
|
||||
pub async fn on_ack(&self, payload: Bytes) {
|
||||
let _ = self.ack_tx.send(payload).await;
|
||||
pub fn on_ack(&self, payload: Bytes) {
|
||||
let _ = self.ack_tx.try_send(payload);
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for HeartbeatHandle {
|
||||
fn drop(&mut self) {
|
||||
if let Some(task) = self.task.take() {
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,7 +56,7 @@ impl HeartbeatHandle {
|
||||
pub fn spawn_noop() -> HeartbeatHandle {
|
||||
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
||||
// receiver is immediately dropped; on_ack() calls will silently fail
|
||||
HeartbeatHandle { ack_tx }
|
||||
HeartbeatHandle { ack_tx, task: None }
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
@@ -74,7 +82,7 @@ pub fn spawn(
|
||||
) -> HeartbeatHandle {
|
||||
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
||||
|
||||
tokio::spawn(async move {
|
||||
let task = tokio::spawn(async move {
|
||||
// Read initial interval from dynamic config (may be updated by remote config).
|
||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||
let mut current_interval = initial_interval;
|
||||
@@ -151,7 +159,8 @@ pub fn spawn(
|
||||
current_interval = new_interval;
|
||||
}
|
||||
}
|
||||
Some(ack_payload) = ack_rx.recv() => {
|
||||
ack_payload = ack_rx.recv() => {
|
||||
let Some(ack_payload) = ack_payload else { break; };
|
||||
match handle_ack(&server, &ack_payload) {
|
||||
AckDecision::Accept {
|
||||
heartbeat_id: ack_id,
|
||||
@@ -179,7 +188,10 @@ pub fn spawn(
|
||||
}
|
||||
});
|
||||
|
||||
HeartbeatHandle { ack_tx }
|
||||
HeartbeatHandle {
|
||||
ack_tx,
|
||||
task: Some(task),
|
||||
}
|
||||
}
|
||||
|
||||
async fn build_heartbeat_payload(
|
||||
|
||||
@@ -3,6 +3,7 @@ pub mod dispatcher;
|
||||
pub mod heartbeat;
|
||||
pub mod protocol;
|
||||
pub mod stream_handler;
|
||||
mod task;
|
||||
pub mod writer;
|
||||
|
||||
use std::sync::Arc;
|
||||
@@ -332,9 +333,9 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
||||
gateway_handle.abort();
|
||||
let _ = (&mut gateway_handle).await;
|
||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||
|
||||
let (_restarted_gateway_state, restarted_gateway_handle) =
|
||||
start_gateway_on_port_retry(gateway_port)
|
||||
@@ -349,6 +350,7 @@ mod tests {
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(server.tunnel_metrics.snapshot().connect_successes >= 2);
|
||||
let _ = shutdown_tx.send(true);
|
||||
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
||||
.await
|
||||
@@ -380,7 +382,17 @@ mod tests {
|
||||
gateway_base_url: &str,
|
||||
node_id: &str,
|
||||
) -> Option<(StatusCode, String)> {
|
||||
let payload = relay_probe_envelope();
|
||||
let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?;
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Some((status, body))
|
||||
}
|
||||
|
||||
async fn relay_response(
|
||||
gateway_base_url: &str,
|
||||
node_id: &str,
|
||||
payload: Vec<u8>,
|
||||
) -> Option<reqwest::Response> {
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.expect("test clock should be after epoch")
|
||||
@@ -398,7 +410,7 @@ mod tests {
|
||||
&nonce,
|
||||
&digest,
|
||||
);
|
||||
let response = reqwest::Client::new()
|
||||
reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
||||
))
|
||||
@@ -421,10 +433,7 @@ mod tests {
|
||||
.body(payload)
|
||||
.send()
|
||||
.await
|
||||
.ok()?;
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Some((status, body))
|
||||
.ok()
|
||||
}
|
||||
|
||||
fn relay_probe_envelope() -> Vec<u8> {
|
||||
@@ -456,30 +465,150 @@ mod tests {
|
||||
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
||||
// The embedded gateway now fails closed when relay authentication is
|
||||
// not configured. Keep this integration fixture explicitly authenticated.
|
||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||
std::env::set_var(
|
||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||
);
|
||||
std::env::set_var(
|
||||
"AETHER_GATEWAY_INSTANCE_ID",
|
||||
"tunnel-reconnect-test-gateway",
|
||||
);
|
||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||
aether_gateway::configure_test_tunnel_security(
|
||||
&mut state,
|
||||
"node-recovery",
|
||||
"test-generation-1",
|
||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||
);
|
||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
|
||||
let state = {
|
||||
let _guard = ENV_LOCK.lock().unwrap();
|
||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||
std::env::set_var(
|
||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||
);
|
||||
std::env::set_var(
|
||||
"AETHER_GATEWAY_INSTANCE_ID",
|
||||
"tunnel-reconnect-test-gateway",
|
||||
);
|
||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||
aether_gateway::configure_test_tunnel_security(
|
||||
&mut state,
|
||||
"node-recovery",
|
||||
"test-generation-1",
|
||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||
);
|
||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||
state
|
||||
};
|
||||
let router = build_router_with_state(state.clone());
|
||||
let handle = spawn_router_on_port(port, router).await?;
|
||||
Ok((state, handle))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() {
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::routing::get;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
ensure_rustls_provider();
|
||||
let upstream_port = reserve_local_port().unwrap();
|
||||
let upstream = Router::new()
|
||||
.route(
|
||||
"/large",
|
||||
get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }),
|
||||
)
|
||||
.route(
|
||||
"/idle",
|
||||
get(|| async {
|
||||
let first = futures_util::stream::once(async {
|
||||
Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n"))
|
||||
});
|
||||
(
|
||||
[("content-type", "text/event-stream")],
|
||||
Body::from_stream(first.chain(futures_util::stream::pending())),
|
||||
)
|
||||
}),
|
||||
);
|
||||
let upstream_task = super::task::SessionTask::new(
|
||||
spawn_router_on_port(upstream_port, upstream).await.unwrap(),
|
||||
);
|
||||
let gateway_port = reserve_local_port().unwrap();
|
||||
let gateway_url = format!("http://127.0.0.1:{gateway_port}");
|
||||
let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap();
|
||||
let gateway_task = super::task::SessionTask::new(gateway_task);
|
||||
let mut config = sample_config(&gateway_url);
|
||||
config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired;
|
||||
config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into());
|
||||
config.tunnel_stream_initial_window_bytes = 512 * 1024;
|
||||
config.tunnel_drain_deadline_ms = 100;
|
||||
config.allow_private_targets = true;
|
||||
config.allowed_ports.push(upstream_port);
|
||||
let state = sample_state(config);
|
||||
let server = sample_server(&state, "node-recovery");
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
let (_drain_tx, drain_rx) = watch::channel(false);
|
||||
let tunnel_task = super::task::SessionTask::new(tokio::spawn({
|
||||
let state = Arc::clone(&state);
|
||||
let server = Arc::clone(&server);
|
||||
async move {
|
||||
run(&state, &server, 0, shutdown_rx, drain_rx).await;
|
||||
}
|
||||
}));
|
||||
wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await;
|
||||
|
||||
let envelope = |path: &str| {
|
||||
let mut meta: protocol::RequestMeta =
|
||||
serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap();
|
||||
meta.url = format!("http://127.0.0.1:{upstream_port}/{path}");
|
||||
meta.stream = true;
|
||||
meta.timeout = 10;
|
||||
meta.stream_first_byte_timeout_ms = Some(10_000);
|
||||
let encoded = serde_json::to_vec(&meta).unwrap();
|
||||
let mut result = (encoded.len() as u32).to_be_bytes().to_vec();
|
||||
result.extend(encoded);
|
||||
result
|
||||
};
|
||||
let response = relay_response(&gateway_url, "node-recovery", envelope("large"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let body = tokio::time::timeout(Duration::from_secs(10), response.bytes())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(body.len(), 2 * 1024 * 1024);
|
||||
assert!(body.iter().all(|byte| *byte == b'x'));
|
||||
|
||||
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
response.chunk().await.unwrap().unwrap(),
|
||||
"data: started\n\n"
|
||||
);
|
||||
drop(response);
|
||||
tokio::time::timeout(Duration::from_secs(3), async {
|
||||
while server
|
||||
.active_connections
|
||||
.load(std::sync::atomic::Ordering::Acquire)
|
||||
!= 0
|
||||
{
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("cancelled SSE must release the upstream handler");
|
||||
|
||||
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(response.chunk().await.unwrap().is_some());
|
||||
shutdown_tx.send(true).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(3), tunnel_task)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
server
|
||||
.active_connections
|
||||
.load(std::sync::atomic::Ordering::Acquire),
|
||||
0
|
||||
);
|
||||
drop(response);
|
||||
drop(gateway_task);
|
||||
drop(upstream_task);
|
||||
}
|
||||
|
||||
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
||||
if let Some(value) = value {
|
||||
std::env::set_var(key, value);
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use tokio::task::{JoinError, JoinHandle};
|
||||
|
||||
pub(super) struct SessionTask<T>(JoinHandle<T>);
|
||||
|
||||
impl<T> SessionTask<T> {
|
||||
pub(super) fn new(handle: JoinHandle<T>) -> Self {
|
||||
Self(handle)
|
||||
}
|
||||
pub(super) fn abort(&self) {
|
||||
self.0.abort();
|
||||
}
|
||||
pub(super) fn is_finished(&self) -> bool {
|
||||
self.0.is_finished()
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> Future for SessionTask<T> {
|
||||
type Output = Result<T, JoinError>;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
Pin::new(&mut self.0).poll(context)
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> Drop for SessionTask<T> {
|
||||
fn drop(&mut self) {
|
||||
self.0.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_a_session_task_aborts_its_child() {
|
||||
let child = tokio::spawn(std::future::pending::<()>());
|
||||
let abort = child.abort_handle();
|
||||
drop(SessionTask::new(child));
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), async {
|
||||
while !abort.is_finished() {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -13,6 +13,7 @@ use aether_contracts::tunnel::{MsgType, HEADER_SIZE};
|
||||
use aether_runtime::QueueSnapshot;
|
||||
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
||||
use futures_util::SinkExt;
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, trace};
|
||||
@@ -24,6 +25,8 @@ use aether_contracts::tunnel_security::SecureFrameCodec;
|
||||
|
||||
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
||||
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
||||
const WRITE_TIMEOUT: Duration = Duration::from_secs(15);
|
||||
const CLOSE_TIMEOUT: Duration = Duration::from_secs(1);
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum FramePriority {
|
||||
@@ -43,9 +46,18 @@ pub struct FrameQueueSnapshots {
|
||||
pub struct FrameSender {
|
||||
high_tx: BoundedQueueSender<Frame>,
|
||||
normal_tx: BoundedQueueSender<Frame>,
|
||||
close_tx: watch::Sender<bool>,
|
||||
}
|
||||
|
||||
impl FrameSender {
|
||||
pub fn close(&self) {
|
||||
let _ = self.close_tx.send(true);
|
||||
}
|
||||
|
||||
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||
self.close_tx.subscribe()
|
||||
}
|
||||
|
||||
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
||||
match classify_frame_priority(&frame) {
|
||||
FramePriority::High => self.high_tx.send(frame).await,
|
||||
@@ -73,7 +85,12 @@ impl FrameSender {
|
||||
high_tx: BoundedQueueSender<Frame>,
|
||||
normal_tx: BoundedQueueSender<Frame>,
|
||||
) -> Self {
|
||||
Self { high_tx, normal_tx }
|
||||
let (close_tx, _) = watch::channel(false);
|
||||
Self {
|
||||
high_tx,
|
||||
normal_tx,
|
||||
close_tx,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -113,15 +130,24 @@ where
|
||||
{
|
||||
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
||||
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
||||
let tx = FrameSender { high_tx, normal_tx };
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
let tx = FrameSender {
|
||||
high_tx,
|
||||
normal_tx,
|
||||
close_tx,
|
||||
};
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
let mut ping_ticker = tokio::time::interval(ping_interval);
|
||||
let mut high_open = true;
|
||||
let mut normal_open = true;
|
||||
let mut close_open = true;
|
||||
ping_ticker.tick().await; // skip first immediate tick
|
||||
|
||||
loop {
|
||||
if *close_rx.borrow() {
|
||||
break;
|
||||
}
|
||||
if let Ok(frame) = high_rx.try_recv() {
|
||||
if !write_frame(
|
||||
&mut sink,
|
||||
@@ -141,6 +167,10 @@ where
|
||||
|
||||
tokio::select! {
|
||||
biased;
|
||||
changed = close_rx.changed(), if close_open => {
|
||||
if changed.is_err() { close_open = false; }
|
||||
if *close_rx.borrow() { break; }
|
||||
},
|
||||
frame = high_rx.recv(), if high_open => {
|
||||
match frame {
|
||||
Some(frame) => {
|
||||
@@ -152,7 +182,7 @@ where
|
||||
}
|
||||
}
|
||||
_ = ping_ticker.tick(), if high_open || normal_open => {
|
||||
if let Err(e) = sink.send(Message::Ping(vec![])).await {
|
||||
if let Err(e) = send_message(&mut sink, Message::Ping(vec![])).await {
|
||||
error!(error = %e, "failed to send WebSocket ping");
|
||||
if let Some(metrics) = tunnel_metrics.as_deref() {
|
||||
metrics.record_error("ws_ping_error", &e.to_string());
|
||||
@@ -174,7 +204,7 @@ where
|
||||
}
|
||||
}
|
||||
debug!("writer task exiting");
|
||||
let _ = sink.close().await;
|
||||
let _ = tokio::time::timeout(CLOSE_TIMEOUT, sink.close()).await;
|
||||
});
|
||||
|
||||
(tx, handle)
|
||||
@@ -228,7 +258,7 @@ where
|
||||
None => frame.encode(),
|
||||
};
|
||||
let wire_len = data.len().max(HEADER_SIZE);
|
||||
if let Err(e) = sink.send(Message::Binary(data.into())).await {
|
||||
if let Err(e) = send_message(sink, Message::Binary(data.into())).await {
|
||||
error!(
|
||||
stream_id = stream_id,
|
||||
msg_type = ?msg_type,
|
||||
@@ -248,8 +278,92 @@ where
|
||||
true
|
||||
}
|
||||
|
||||
async fn send_message<S>(
|
||||
sink: &mut S,
|
||||
message: Message,
|
||||
) -> Result<(), tokio_tungstenite::tungstenite::Error>
|
||||
where
|
||||
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
|
||||
{
|
||||
tokio::time::timeout(WRITE_TIMEOUT, sink.send(message))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
tokio_tungstenite::tungstenite::Error::Io(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
"tunnel WebSocket write timed out",
|
||||
))
|
||||
})?
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[tokio::test]
|
||||
async fn dropping_last_sender_flushes_queued_body_and_end_frames() {
|
||||
let sink = VecSink::default();
|
||||
let sent = Arc::clone(&sink.sent);
|
||||
let (sender, task) = spawn_writer(sink, Duration::from_secs(60));
|
||||
sender
|
||||
.send(Frame::new(
|
||||
7,
|
||||
MsgType::ResponseBody,
|
||||
0,
|
||||
bytes::Bytes::from_static(b"late"),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
sender
|
||||
.send(Frame::new(7, MsgType::StreamEnd, 0, bytes::Bytes::new()))
|
||||
.await
|
||||
.unwrap();
|
||||
drop(sender);
|
||||
task.await.unwrap();
|
||||
let frames = sent.lock().unwrap();
|
||||
assert_eq!(frames.len(), 2);
|
||||
let Message::Binary(body) = &frames[0] else {
|
||||
panic!("expected body")
|
||||
};
|
||||
assert_eq!(
|
||||
Frame::decode(body.clone().into()).unwrap().payload,
|
||||
b"late".as_slice()
|
||||
);
|
||||
}
|
||||
|
||||
struct StalledSink;
|
||||
|
||||
impl futures_util::Sink<Message> for StalledSink {
|
||||
type Error = Error;
|
||||
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Error> {
|
||||
Ok(())
|
||||
}
|
||||
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||
Poll::Pending
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test(start_paused = true)]
|
||||
async fn stalled_socket_write_and_close_are_bounded() {
|
||||
let (sender, task) = spawn_writer(StalledSink, Duration::from_secs(60));
|
||||
sender
|
||||
.send(Frame::new(
|
||||
1,
|
||||
MsgType::ResponseBody,
|
||||
0,
|
||||
bytes::Bytes::from_static(b"data"),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(20), task)
|
||||
.await
|
||||
.expect("writer should time out")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
Reference in New Issue
Block a user