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:
@@ -0,0 +1,186 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::Arc;
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use parking_lot::Mutex;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
const CHUNK_BYTES: usize = 32 * 1024;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LocalBodyEvent {
|
||||
Chunk(Bytes),
|
||||
End,
|
||||
Error(String),
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct BufferState {
|
||||
chunks: VecDeque<BytesMut>,
|
||||
bytes: usize,
|
||||
terminal: Option<Result<(), String>>,
|
||||
receiver_taken: bool,
|
||||
receiver_closed: bool,
|
||||
}
|
||||
|
||||
pub(super) struct ResponseBuffer {
|
||||
state: Mutex<BufferState>,
|
||||
notify: Notify,
|
||||
capacity: usize,
|
||||
}
|
||||
|
||||
impl ResponseBuffer {
|
||||
pub(super) fn new(capacity: usize) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
state: Mutex::new(BufferState::default()),
|
||||
notify: Notify::new(),
|
||||
capacity,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
|
||||
let mut state = self.state.lock();
|
||||
if state.receiver_taken {
|
||||
return None;
|
||||
}
|
||||
state.receiver_taken = true;
|
||||
Some(BodyReceiver {
|
||||
buffer: Arc::clone(self),
|
||||
finished: false,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn push(&self, mut payload: Bytes) -> bool {
|
||||
let mut state = self.state.lock();
|
||||
if state.terminal.is_some()
|
||||
|| state.receiver_closed
|
||||
|| payload.len() > self.capacity.saturating_sub(state.bytes)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
state.bytes += payload.len();
|
||||
while !payload.is_empty() {
|
||||
if let Some(tail) = state
|
||||
.chunks
|
||||
.back_mut()
|
||||
.filter(|chunk| chunk.len() < CHUNK_BYTES)
|
||||
{
|
||||
let count = payload.len().min(CHUNK_BYTES - tail.len());
|
||||
tail.extend_from_slice(&payload.split_to(count));
|
||||
} else {
|
||||
let count = payload.len().min(CHUNK_BYTES);
|
||||
let chunk = payload.split_to(count);
|
||||
state.chunks.push_back(
|
||||
chunk
|
||||
.try_into_mut()
|
||||
.unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())),
|
||||
);
|
||||
}
|
||||
}
|
||||
drop(state);
|
||||
self.notify.notify_waiters();
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn finish(&self, result: Result<(), String>) {
|
||||
let mut state = self.state.lock();
|
||||
if state.terminal.is_none() {
|
||||
state.terminal = Some(result);
|
||||
}
|
||||
drop(state);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct BodyReceiver {
|
||||
buffer: Arc<ResponseBuffer>,
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl BodyReceiver {
|
||||
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||
if self.finished {
|
||||
return None;
|
||||
}
|
||||
loop {
|
||||
let notified = self.buffer.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
{
|
||||
let mut state = self.buffer.state.lock();
|
||||
if let Some(chunk) = state.chunks.pop_front() {
|
||||
state.bytes -= chunk.len();
|
||||
return Some(LocalBodyEvent::Chunk(chunk.freeze()));
|
||||
}
|
||||
if let Some(terminal) = state.terminal.take() {
|
||||
self.finished = true;
|
||||
state.receiver_closed = true;
|
||||
return Some(match terminal {
|
||||
Ok(()) => LocalBodyEvent::End,
|
||||
Err(error) => LocalBodyEvent::Error(error),
|
||||
});
|
||||
}
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for BodyReceiver {
|
||||
fn drop(&mut self) {
|
||||
let mut state = self.buffer.state.lock();
|
||||
state.receiver_closed = true;
|
||||
state.chunks.clear();
|
||||
state.bytes = 0;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn error_survives_a_full_buffer() {
|
||||
let buffer = ResponseBuffer::new(CHUNK_BYTES);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES])));
|
||||
buffer.finish(Err("proxy disconnected".into()));
|
||||
assert!(matches!(
|
||||
receiver.recv().await,
|
||||
Some(LocalBodyEvent::Chunk(_))
|
||||
));
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected")
|
||||
);
|
||||
assert!(receiver.recv().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn small_frames_are_coalesced_within_the_byte_budget() {
|
||||
let buffer = ResponseBuffer::new(4096);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
for _ in 0..4096 {
|
||||
assert!(buffer.push(Bytes::from_static(b"x")));
|
||||
}
|
||||
assert!(!buffer.push(Bytes::from_static(b"x")));
|
||||
assert_eq!(buffer.state.lock().chunks.len(), 1);
|
||||
buffer.finish(Ok(()));
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096)
|
||||
);
|
||||
assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() {
|
||||
let buffer = ResponseBuffer::new(1024);
|
||||
let mut receiver = buffer.take_receiver().unwrap();
|
||||
let task = tokio::spawn(async move { receiver.recv().await });
|
||||
tokio::task::yield_now().await;
|
||||
buffer.finish(Err("cancelled".into()));
|
||||
buffer.finish(Ok(()));
|
||||
assert!(
|
||||
matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled")
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,250 @@
|
||||
use super::*;
|
||||
|
||||
async fn fixture(
|
||||
window: u32,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Arc<HubRouter>,
|
||||
Arc<ProxyConn>,
|
||||
Arc<LocalStream>,
|
||||
aether_runtime::BoundedQueueReceiver<Message>,
|
||||
) {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
let (sender, receiver) = bounded_queue(capacity);
|
||||
let (close_tx, _) = watch::channel(false);
|
||||
let connection = Arc::new(
|
||||
ProxyConn::new(
|
||||
99,
|
||||
"flow-test".into(),
|
||||
"flow-test".into(),
|
||||
sender,
|
||||
close_tx,
|
||||
16,
|
||||
3,
|
||||
)
|
||||
.with_settings(protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: window,
|
||||
min_window_update_bytes: (window / 4).max(1),
|
||||
drain_deadline_ms: 1000,
|
||||
}),
|
||||
);
|
||||
hub.register_proxy(Arc::clone(&connection));
|
||||
let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
(hub, connection, stream, receiver)
|
||||
}
|
||||
|
||||
fn meta() -> protocol::RequestMeta {
|
||||
protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "GET".into(),
|
||||
url: "https://example.com".into(),
|
||||
headers: HashMap::new(),
|
||||
stream: true,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
}
|
||||
}
|
||||
|
||||
async fn headers(hub: &Arc<HubRouter>, stream: &LocalStream) {
|
||||
let payload = serde_json::to_vec(&protocol::ResponseMeta {
|
||||
status: 200,
|
||||
headers: vec![],
|
||||
})
|
||||
.unwrap();
|
||||
let mut frame = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_HEADERS,
|
||||
0,
|
||||
&payload,
|
||||
);
|
||||
hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() {
|
||||
let (hub, _, stream, mut outbound) = fixture(128, 1).await;
|
||||
headers(&hub, &stream).await;
|
||||
assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64])));
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
assert!(
|
||||
tokio::time::timeout(Duration::from_millis(10), receiver.recv())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert_eq!(*stream.response_consumed_since_update.lock(), 64);
|
||||
outbound.recv().await.unwrap();
|
||||
let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64));
|
||||
assert_eq!(*stream.response_consumed_since_update.lock(), 0);
|
||||
let Message::Binary(data) = outbound.recv().await.unwrap() else {
|
||||
panic!("expected binary update")
|
||||
};
|
||||
let frame = aether_contracts::tunnel::Frame::decode(data).unwrap();
|
||||
let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap();
|
||||
assert_eq!(
|
||||
frame.msg_type,
|
||||
aether_contracts::tunnel::MsgType::WindowUpdate
|
||||
);
|
||||
assert_eq!(update.delta_bytes, 64);
|
||||
hub.cancel_local_stream(stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn response_credit_is_not_returned_until_consumed() {
|
||||
let (hub, _, stream, mut outbound) = fixture(128, 4).await;
|
||||
outbound.recv().await.unwrap();
|
||||
let mut body = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
&[b'x'; 128],
|
||||
);
|
||||
hub.handle_proxy_frame(99, &mut body).await;
|
||||
assert!(outbound.try_recv().is_err());
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
assert!(
|
||||
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128)
|
||||
);
|
||||
assert!(outbound.try_recv().is_ok());
|
||||
hub.cancel_local_stream(stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_stream_open_releases_slot_without_resetting_connection() {
|
||||
let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await;
|
||||
let opening_hub = Arc::clone(&hub);
|
||||
let opening =
|
||||
tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await });
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while hub.local_streams.len() != 2 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
opening.abort();
|
||||
assert!(matches!(opening.await, Err(error) if error.is_cancelled()));
|
||||
assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1);
|
||||
assert_eq!(hub.local_streams.len(), 1);
|
||||
assert_eq!(hub.proxy_to_local.len(), 1);
|
||||
assert!(connection.is_available());
|
||||
outbound.recv().await.unwrap();
|
||||
assert!(outbound.try_recv().is_err());
|
||||
let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
outbound.recv().await.unwrap();
|
||||
hub.cancel_local_stream(first_stream.id, "test complete");
|
||||
outbound.recv().await.unwrap();
|
||||
hub.cancel_local_stream(next_stream.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn full_response_buffer_preserves_disconnect_error() {
|
||||
let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await;
|
||||
outbound.recv().await.unwrap();
|
||||
headers(&hub, &stream).await;
|
||||
let mut receiver = stream.take_body_receiver().unwrap();
|
||||
for _ in 0..128 {
|
||||
let mut frame = protocol::encode_frame(
|
||||
stream.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
&vec![b'x'; 32 * 1024],
|
||||
);
|
||||
hub.handle_proxy_frame(99, &mut frame).await;
|
||||
}
|
||||
hub.unregister_proxy(connection.id, &connection.node_id);
|
||||
let mut bytes = 0;
|
||||
loop {
|
||||
match receiver.recv().await {
|
||||
Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(),
|
||||
Some(LocalBodyEvent::Error(error)) => {
|
||||
assert!(error.contains("disconnected"));
|
||||
break;
|
||||
}
|
||||
event => panic!("disconnect must not become normal EOF: {event:?}"),
|
||||
}
|
||||
}
|
||||
assert_eq!(bytes, 4 * 1024 * 1024);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn slow_stream_does_not_block_another_stream_on_the_same_connection() {
|
||||
let (hub, _, slow, mut outbound) = fixture(128, 512).await;
|
||||
outbound.recv().await.unwrap();
|
||||
let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||
assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128])));
|
||||
let mut overflowing = protocol::encode_frame(
|
||||
slow.proxy_stream_id,
|
||||
protocol::RESPONSE_BODY,
|
||||
0,
|
||||
b"overflow",
|
||||
);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
hub.handle_proxy_frame(99, &mut overflowing).await;
|
||||
headers(&hub, &fast).await;
|
||||
assert_eq!(
|
||||
fast.wait_headers(Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap()
|
||||
.status,
|
||||
200
|
||||
);
|
||||
})
|
||||
.await
|
||||
.expect("slow stream must not block connection reader");
|
||||
assert!(!hub.local_streams.contains_key(&slow.id));
|
||||
assert!(hub.local_streams.contains_key(&fast.id));
|
||||
hub.cancel_local_stream(fast.id, "test complete");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelling_a_stream_wakes_request_window_waiters() {
|
||||
let (_, _, stream, _) = fixture(128, 512).await;
|
||||
*stream.request_window.available.lock() = 0;
|
||||
let waiter = tokio::spawn({
|
||||
let stream = Arc::clone(&stream);
|
||||
async move {
|
||||
stream
|
||||
.acquire_request_window(1, Duration::from_secs(30))
|
||||
.await
|
||||
}
|
||||
});
|
||||
tokio::task::yield_now().await;
|
||||
stream.fail("cancelled");
|
||||
assert!(tokio::time::timeout(Duration::from_secs(1), waiter)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() {
|
||||
for index in 0..256 {
|
||||
let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1));
|
||||
let window = Arc::new(StreamFlowWindow::new(0));
|
||||
let waiter = tokio::spawn({
|
||||
let stream = Arc::clone(&stream);
|
||||
let window = Arc::clone(&window);
|
||||
async move {
|
||||
stream.wait_headers(Duration::from_secs(1)).await.unwrap();
|
||||
window.acquire(1, Duration::from_secs(1)).await.unwrap();
|
||||
}
|
||||
});
|
||||
stream.set_response_headers(protocol::ResponseMeta {
|
||||
status: 200,
|
||||
headers: vec![],
|
||||
});
|
||||
window.add(1);
|
||||
waiter.await.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -12,10 +12,11 @@ use axum::extract::ws::Message;
|
||||
use bytes::Bytes;
|
||||
use dashmap::DashMap;
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::sync::{watch, Notify};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
pub use super::body::LocalBodyEvent;
|
||||
use super::body::{BodyReceiver, ResponseBuffer};
|
||||
use super::control_plane::ControlPlaneClient;
|
||||
use super::protocol;
|
||||
|
||||
@@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
|
||||
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
||||
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "flow_control_tests.rs"]
|
||||
mod flow_control_tests;
|
||||
|
||||
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
||||
.ok()
|
||||
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
|
||||
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
||||
});
|
||||
|
||||
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||
STREAM_INITIAL_WINDOW_BYTES
|
||||
.saturating_div(4)
|
||||
.clamp(1, 1024 * 1024)
|
||||
});
|
||||
pub(super) fn local_settings() -> protocol::SettingsPayload {
|
||||
protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES)
|
||||
.min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32),
|
||||
min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES
|
||||
.saturating_div(4)
|
||||
.clamp(1, 1024 * 1024),
|
||||
drain_deadline_ms: *DRAIN_DEADLINE_MS,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SendStatus {
|
||||
@@ -91,6 +101,7 @@ impl ConnHealthState {
|
||||
struct StreamFlowWindow {
|
||||
available: Mutex<u64>,
|
||||
notify: Notify,
|
||||
closed: AtomicBool,
|
||||
}
|
||||
|
||||
impl StreamFlowWindow {
|
||||
@@ -98,6 +109,7 @@ impl StreamFlowWindow {
|
||||
Self {
|
||||
available: Mutex::new(u64::from(initial)),
|
||||
notify: Notify::new(),
|
||||
closed: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,6 +121,12 @@ impl StreamFlowWindow {
|
||||
let requested = bytes as u64;
|
||||
let started_at = Instant::now();
|
||||
loop {
|
||||
let notified = self.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
if self.closed.load(Ordering::Acquire) {
|
||||
return Err(());
|
||||
}
|
||||
{
|
||||
let mut available = self.available.lock();
|
||||
if *available >= requested {
|
||||
@@ -120,10 +138,7 @@ impl StreamFlowWindow {
|
||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||
return Err(());
|
||||
};
|
||||
if tokio::time::timeout(remaining, self.notify.notified())
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||
return Err(());
|
||||
}
|
||||
}
|
||||
@@ -138,6 +153,11 @@ impl StreamFlowWindow {
|
||||
drop(available);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
|
||||
fn close(&self) {
|
||||
self.closed.store(true, Ordering::Release);
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
@@ -209,6 +229,10 @@ impl BoundedOutbound {
|
||||
pub fn snapshot(&self) -> QueueSnapshot {
|
||||
self.tx.snapshot()
|
||||
}
|
||||
|
||||
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||
self.close_tx.subscribe()
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ProxyConn {
|
||||
@@ -231,6 +255,7 @@ pub struct ProxyConn {
|
||||
flow_window_blocked_ms: AtomicU64,
|
||||
write_latency_last_us: AtomicU64,
|
||||
write_latency_ewma_us: AtomicU64,
|
||||
settings: Mutex<protocol::SettingsPayload>,
|
||||
}
|
||||
|
||||
impl ProxyConn {
|
||||
@@ -244,6 +269,7 @@ impl ProxyConn {
|
||||
protocol_version: u8,
|
||||
) -> Self {
|
||||
Self {
|
||||
settings: Mutex::new(local_settings()),
|
||||
id,
|
||||
node_id,
|
||||
node_name,
|
||||
@@ -271,6 +297,11 @@ impl ProxyConn {
|
||||
self
|
||||
}
|
||||
|
||||
pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self {
|
||||
*self.settings.get_mut() = settings;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||
self.node_generation = tunnel_generation;
|
||||
self
|
||||
@@ -565,13 +596,6 @@ pub struct LocalResponseHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LocalBodyEvent {
|
||||
Chunk(Bytes),
|
||||
End,
|
||||
Error(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct LocalWaitState {
|
||||
response: Option<LocalResponseHead>,
|
||||
@@ -585,10 +609,11 @@ pub struct LocalStream {
|
||||
proxy_stream_id: u32,
|
||||
request_window: StreamFlowWindow,
|
||||
response_consumed_since_update: Mutex<u64>,
|
||||
min_window_update_bytes: u32,
|
||||
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
|
||||
wait_state: Mutex<LocalWaitState>,
|
||||
headers_notify: Notify,
|
||||
body_tx: mpsc::Sender<LocalBodyEvent>,
|
||||
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
|
||||
body: Arc<ResponseBuffer>,
|
||||
terminal: AtomicBool,
|
||||
}
|
||||
|
||||
@@ -600,7 +625,6 @@ impl LocalStream {
|
||||
proxy_stream_id: u32,
|
||||
initial_window_bytes: u32,
|
||||
) -> Self {
|
||||
let (body_tx, body_rx) = mpsc::channel(128);
|
||||
Self {
|
||||
id,
|
||||
tunnel_generation,
|
||||
@@ -608,10 +632,11 @@ impl LocalStream {
|
||||
proxy_stream_id,
|
||||
request_window: StreamFlowWindow::new(initial_window_bytes),
|
||||
response_consumed_since_update: Mutex::new(0),
|
||||
min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024),
|
||||
response_connection: Mutex::new(None),
|
||||
wait_state: Mutex::new(LocalWaitState::default()),
|
||||
headers_notify: Notify::new(),
|
||||
body_tx,
|
||||
body_rx: Mutex::new(Some(body_rx)),
|
||||
body: ResponseBuffer::new(initial_window_bytes as usize),
|
||||
terminal: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
@@ -632,26 +657,51 @@ impl LocalStream {
|
||||
self.request_window.add(delta);
|
||||
}
|
||||
|
||||
fn response_window_update_delta(&self, bytes: usize) -> Option<u32> {
|
||||
if bytes == 0 {
|
||||
return None;
|
||||
async fn flush_response_credit(&self) -> Result<(), String> {
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let mut consumed = self.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_add(bytes as u64);
|
||||
let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES);
|
||||
if *consumed < threshold {
|
||||
return None;
|
||||
let connection = self
|
||||
.response_connection
|
||||
.lock()
|
||||
.as_ref()
|
||||
.and_then(std::sync::Weak::upgrade);
|
||||
let Some(connection) = connection else {
|
||||
return Ok(());
|
||||
};
|
||||
if connection.protocol_version() < 3 {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let delta = (*consumed).min(u64::from(u32::MAX)) as u32;
|
||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||
Some(delta)
|
||||
let delta = {
|
||||
let consumed = self.response_consumed_since_update.lock();
|
||||
if *consumed < u64::from(self.min_window_update_bytes) {
|
||||
return Ok(());
|
||||
}
|
||||
(*consumed).min(u64::from(u32::MAX)) as u32
|
||||
};
|
||||
let frame = protocol::encode_window_update(self.proxy_stream_id, delta);
|
||||
if connection
|
||||
.send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT)
|
||||
.await
|
||||
== SendStatus::Queued
|
||||
{
|
||||
let mut consumed = self.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||
return Ok(());
|
||||
}
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return Ok(());
|
||||
}
|
||||
connection.request_close();
|
||||
Err("proxy flow-control update failed".to_string())
|
||||
}
|
||||
|
||||
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
||||
tokio::time::timeout(timeout, async {
|
||||
loop {
|
||||
let notified = self.headers_notify.notified();
|
||||
tokio::pin!(notified);
|
||||
notified.as_mut().enable();
|
||||
let outcome = {
|
||||
let state = self.wait_state.lock();
|
||||
if let Some(response) = &state.response {
|
||||
@@ -662,15 +712,20 @@ impl LocalStream {
|
||||
if let Some(error) = outcome {
|
||||
return Err(error);
|
||||
}
|
||||
self.headers_notify.notified().await;
|
||||
notified.await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map_err(|_| "timed out waiting for response headers".to_string())?
|
||||
}
|
||||
|
||||
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
|
||||
self.body_rx.lock().take()
|
||||
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
|
||||
self.body.take_receiver().map(|receiver| LocalBodyReceiver {
|
||||
receiver,
|
||||
stream: Arc::clone(self),
|
||||
failed: false,
|
||||
pending: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
||||
@@ -690,21 +745,11 @@ impl LocalStream {
|
||||
}
|
||||
}
|
||||
|
||||
async fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||
fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return false;
|
||||
}
|
||||
// Use a timeout to prevent a slow consumer from blocking the shared
|
||||
// proxy-connection reader (head-of-line blocking across streams).
|
||||
match tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(())) => true,
|
||||
_ => false,
|
||||
}
|
||||
self.body.push(payload)
|
||||
}
|
||||
|
||||
fn finish(&self) {
|
||||
@@ -722,7 +767,8 @@ impl LocalStream {
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::End);
|
||||
self.request_window.close();
|
||||
self.body.finish(Ok(()));
|
||||
}
|
||||
|
||||
fn fail(&self, error: impl Into<String>) {
|
||||
@@ -742,7 +788,38 @@ impl LocalStream {
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
|
||||
self.request_window.close();
|
||||
self.body.finish(Err(error));
|
||||
}
|
||||
}
|
||||
|
||||
pub struct LocalBodyReceiver {
|
||||
receiver: BodyReceiver,
|
||||
stream: Arc<LocalStream>,
|
||||
failed: bool,
|
||||
pending: Option<LocalBodyEvent>,
|
||||
}
|
||||
|
||||
impl LocalBodyReceiver {
|
||||
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||
if self.failed {
|
||||
return None;
|
||||
}
|
||||
if self.pending.is_none() {
|
||||
let event = self.receiver.recv().await?;
|
||||
if let LocalBodyEvent::Chunk(chunk) = &event {
|
||||
let mut consumed = self.stream.response_consumed_since_update.lock();
|
||||
*consumed = consumed.saturating_add(chunk.len() as u64);
|
||||
}
|
||||
self.pending = Some(event);
|
||||
}
|
||||
if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) {
|
||||
if let Err(error) = self.stream.flush_response_credit().await {
|
||||
self.failed = true;
|
||||
return Some(LocalBodyEvent::Error(error));
|
||||
}
|
||||
}
|
||||
self.pending.take()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -765,6 +842,21 @@ pub struct HubRouter {
|
||||
drain_reasons: Mutex<HashMap<String, u64>>,
|
||||
}
|
||||
|
||||
struct PendingStreamGuard<'router> {
|
||||
hub: &'router HubRouter,
|
||||
connection: &'router ProxyConn,
|
||||
stream_id: u64,
|
||||
committed: bool,
|
||||
}
|
||||
|
||||
impl Drop for PendingStreamGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
if !self.committed && self.hub.cleanup_local_stream(self.stream_id) {
|
||||
self.connection.release_stream();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct NodeStatusEvent {
|
||||
node_id: String,
|
||||
authenticated_key: Option<String>,
|
||||
@@ -1166,17 +1258,27 @@ impl HubRouter {
|
||||
|
||||
// Frames encoded successfully -- now register the stream.
|
||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||
let local_stream = Arc::new(LocalStream::new(
|
||||
let settings = proxy_conn.settings.lock().clone();
|
||||
let mut local_stream = LocalStream::new(
|
||||
local_stream_id,
|
||||
proxy_conn.node_generation.clone(),
|
||||
proxy_conn.id,
|
||||
proxy_stream_id,
|
||||
*STREAM_INITIAL_WINDOW_BYTES,
|
||||
));
|
||||
settings.initial_stream_window_bytes,
|
||||
);
|
||||
local_stream.min_window_update_bytes = settings.min_window_update_bytes;
|
||||
*local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn));
|
||||
let local_stream = Arc::new(local_stream);
|
||||
self.local_streams
|
||||
.insert(local_stream_id, local_stream.clone());
|
||||
self.proxy_to_local
|
||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||
let mut pending_stream = PendingStreamGuard {
|
||||
hub: self,
|
||||
connection: &proxy_conn,
|
||||
stream_id: local_stream_id,
|
||||
committed: false,
|
||||
};
|
||||
|
||||
let send_status = proxy_conn
|
||||
.send_wait(
|
||||
@@ -1195,10 +1297,11 @@ impl HubRouter {
|
||||
"open_local_stream dispatched"
|
||||
);
|
||||
match send_status {
|
||||
SendStatus::Queued => Ok(local_stream),
|
||||
SendStatus::Queued => {
|
||||
pending_stream.committed = true;
|
||||
Ok(local_stream)
|
||||
}
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
self.cleanup_local_stream(local_stream_id);
|
||||
proxy_conn.release_stream();
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
@@ -1243,7 +1346,9 @@ impl HubRouter {
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||
|
||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
||||
let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE
|
||||
.min(proxy_conn.settings.lock().initial_stream_window_bytes as usize);
|
||||
let total_chunks = payload.len().div_ceil(chunk_size);
|
||||
let result = if total_chunks == 0 {
|
||||
if end_stream {
|
||||
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
||||
@@ -1252,7 +1357,7 @@ impl HubRouter {
|
||||
Ok(())
|
||||
}
|
||||
} else {
|
||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
||||
for (index, chunk) in payload.chunks(chunk_size).enumerate() {
|
||||
let is_last_chunk = index + 1 == total_chunks;
|
||||
if let Err(error) = self
|
||||
.send_request_body_frame(
|
||||
@@ -1352,17 +1457,20 @@ impl HubRouter {
|
||||
} else {
|
||||
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
||||
};
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
if pc.send(Message::Binary(frame.into())) != SendStatus::Queued {
|
||||
pc.request_close();
|
||||
}
|
||||
}
|
||||
stream.fail(reason.to_string());
|
||||
}
|
||||
|
||||
fn cleanup_local_stream(&self, local_stream_id: u64) {
|
||||
fn cleanup_local_stream(&self, local_stream_id: u64) -> bool {
|
||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||
return;
|
||||
return false;
|
||||
};
|
||||
self.proxy_to_local
|
||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||
true
|
||||
}
|
||||
|
||||
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
||||
@@ -1433,9 +1541,7 @@ impl HubRouter {
|
||||
.get(&proxy_conn_id)
|
||||
.map(|entry| entry.value().clone());
|
||||
if let Some(pc) = pc {
|
||||
let _ = pc
|
||||
.send_wait(Message::Binary(pong.into()), Duration::from_millis(250))
|
||||
.await;
|
||||
let _ = pc.send(Message::Binary(pong.into()));
|
||||
}
|
||||
}
|
||||
protocol::PONG => {}
|
||||
@@ -1509,11 +1615,36 @@ impl HubRouter {
|
||||
);
|
||||
}
|
||||
protocol::SETTINGS => {
|
||||
debug!(
|
||||
msg_type = header.msg_type,
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
"received tunnel protocol v3 SETTINGS from proxy"
|
||||
);
|
||||
let settings = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::SettingsPayload>(&payload).ok()
|
||||
})
|
||||
.filter(|settings| settings.is_valid());
|
||||
if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
if header.stream_id != 0 || header.flags != 0 {
|
||||
connection.request_close();
|
||||
return;
|
||||
}
|
||||
let Some(settings) = settings else {
|
||||
connection.request_close();
|
||||
return;
|
||||
};
|
||||
let local = local_settings();
|
||||
let settings = settings
|
||||
.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||
let mut current = connection.settings.lock();
|
||||
if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings {
|
||||
drop(current);
|
||||
connection.request_close();
|
||||
return;
|
||||
}
|
||||
*current = settings;
|
||||
}
|
||||
}
|
||||
protocol::WINDOW_UPDATE => {
|
||||
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
||||
@@ -1739,19 +1870,8 @@ impl HubRouter {
|
||||
None => return,
|
||||
};
|
||||
|
||||
let payload_len = payload.len();
|
||||
if !stream.push_body_chunk(Bytes::from(payload)).await {
|
||||
if !stream.push_body_chunk(Bytes::from(payload)) {
|
||||
self.cancel_local_stream(local_id, "local relay response congested");
|
||||
return;
|
||||
}
|
||||
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
if pc.protocol_version() >= 3 {
|
||||
if let Some(delta) = stream.response_window_update_delta(payload_len) {
|
||||
let frame = protocol::encode_window_update(header.stream_id, delta);
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -9,14 +9,13 @@ use axum::body::{Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::api::response::apply_streaming_response_headers;
|
||||
use crate::headers::should_skip_response_header;
|
||||
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
||||
|
||||
use super::hub::{LocalBodyEvent, LocalStream};
|
||||
use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream};
|
||||
use super::protocol;
|
||||
use super::{AppState, RelayRequestAuthenticated};
|
||||
|
||||
@@ -40,7 +39,7 @@ impl Drop for StreamGuard {
|
||||
pub(crate) struct DirectRelayResponse {
|
||||
status: u16,
|
||||
headers: Vec<(String, String)>,
|
||||
body_rx: mpsc::Receiver<LocalBodyEvent>,
|
||||
body_rx: LocalBodyReceiver,
|
||||
request_guard: StreamGuard,
|
||||
_request_permit: Option<AdmissionPermit>,
|
||||
}
|
||||
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
|
||||
}
|
||||
|
||||
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
||||
if self.request_guard.finished {
|
||||
return Ok(None);
|
||||
}
|
||||
let event = self.body_rx.recv().await;
|
||||
match event {
|
||||
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
||||
Some(LocalBodyEvent::End) | None => {
|
||||
Some(LocalBodyEvent::End) => {
|
||||
self.request_guard.finished = true;
|
||||
Ok(None)
|
||||
}
|
||||
@@ -66,6 +68,7 @@ impl DirectRelayResponse {
|
||||
self.request_guard.finished = true;
|
||||
Err(error)
|
||||
}
|
||||
None => Err("tunnel response ended without a terminal frame".to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
.open_authorized_local_stream(node_id, &meta)
|
||||
.await
|
||||
.map_err(|error| format!("connect: {error}"))?;
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(stream.id, body, true)
|
||||
@@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
status: response_head.status,
|
||||
headers: response_head.headers,
|
||||
body_rx,
|
||||
request_guard: StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
},
|
||||
request_guard,
|
||||
_request_permit: request_permit,
|
||||
})
|
||||
}
|
||||
@@ -259,6 +263,11 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
};
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
let body_stream = match spool.body_stream().await {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
@@ -306,12 +315,6 @@ pub async fn relay_request(
|
||||
);
|
||||
}
|
||||
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
|
||||
let wait_timeout = relay_header_timeout(&meta);
|
||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||
Ok(response) => response,
|
||||
@@ -373,6 +376,9 @@ pub async fn relay_request(
|
||||
}
|
||||
}
|
||||
}
|
||||
if !guard.finished {
|
||||
yield Err(io::Error::other("tunnel response ended without a terminal frame"));
|
||||
}
|
||||
guard.finished = true;
|
||||
};
|
||||
|
||||
@@ -563,6 +569,100 @@ mod tests {
|
||||
request
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancelled_relays_reset_streams_during_upload_and_header_wait() {
|
||||
for direct in [true, false] {
|
||||
for during_upload in [true, false] {
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
sample_connected_proxy_node("node-123"),
|
||||
]));
|
||||
let data = Arc::new(
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(repository)
|
||||
.with_system_config_values_for_tests(
|
||||
Vec::<(String, serde_json::Value)>::new(),
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let state = test_app_state().with_data(data);
|
||||
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let connection = Arc::new(
|
||||
ProxyConn::new(
|
||||
500,
|
||||
"node-123".into(),
|
||||
"Node 123".into(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
3,
|
||||
)
|
||||
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string())
|
||||
.with_settings(protocol::SettingsPayload {
|
||||
initial_stream_window_bytes: 128,
|
||||
min_window_update_bytes: 32,
|
||||
drain_deadline_ms: 1000,
|
||||
}),
|
||||
);
|
||||
state.hub.register_proxy(Arc::clone(&connection));
|
||||
let meta = protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "POST".into(),
|
||||
url: "https://example.com/".into(),
|
||||
headers: HashMap::new(),
|
||||
stream: true,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
};
|
||||
let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]);
|
||||
let relay = tokio::spawn(async move {
|
||||
if direct {
|
||||
let _response =
|
||||
super::open_direct_relay_stream(&state, "node-123", meta, body)
|
||||
.await
|
||||
.unwrap();
|
||||
} else {
|
||||
let request =
|
||||
authenticated_request(encode_relay_envelope(&meta, &body)).await;
|
||||
let _response = relay_request(
|
||||
Path("node-123".into()),
|
||||
State(state),
|
||||
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
|
||||
request,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
});
|
||||
recv_tunnel_test_frame(&mut proxy_rx, "request headers").await;
|
||||
recv_tunnel_test_frame(&mut proxy_rx, "request body").await;
|
||||
relay.abort();
|
||||
assert!(relay.await.unwrap_err().is_cancelled());
|
||||
let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await
|
||||
else {
|
||||
panic!("expected binary reset frame")
|
||||
};
|
||||
let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap();
|
||||
assert_eq!(
|
||||
frame.msg_type,
|
||||
aether_contracts::tunnel::MsgType::ResetStream
|
||||
);
|
||||
assert_eq!(
|
||||
connection
|
||||
.stream_count
|
||||
.load(std::sync::atomic::Ordering::Relaxed),
|
||||
0
|
||||
);
|
||||
assert!(connection.is_available());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
||||
let meta = protocol::RequestMeta {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
mod body;
|
||||
mod control_plane;
|
||||
mod hub;
|
||||
mod local_relay;
|
||||
|
||||
@@ -9,6 +9,7 @@ use aether_runtime::bounded_queue;
|
||||
use axum::extract::ws::{Message, WebSocket};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::JoinSet;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
||||
@@ -84,6 +85,35 @@ pub async fn handle_proxy_connection(
|
||||
|
||||
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
let settings = if protocol_version >= 3 {
|
||||
let Some(settings) = read_proxy_settings(
|
||||
&mut ws_tx,
|
||||
&mut ws_rx,
|
||||
security.as_deref(),
|
||||
protocol_version,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
warn!(conn_id, "proxy SETTINGS negotiation failed");
|
||||
return;
|
||||
};
|
||||
let local = super::hub::local_settings();
|
||||
let negotiated =
|
||||
settings.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||
let message = Message::Binary(protocol::encode_settings(&negotiated).into());
|
||||
let Ok(message) = encrypt_message(message, security.as_deref()) else {
|
||||
return;
|
||||
};
|
||||
if !matches!(
|
||||
tokio::time::timeout(PROXY_HELLO_TIMEOUT, ws_tx.send(message)).await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
return;
|
||||
}
|
||||
negotiated
|
||||
} else {
|
||||
super::hub::local_settings()
|
||||
};
|
||||
let conn = ProxyConn::new(
|
||||
conn_id,
|
||||
node_id.clone(),
|
||||
@@ -93,7 +123,8 @@ pub async fn handle_proxy_connection(
|
||||
max_streams,
|
||||
protocol_version,
|
||||
)
|
||||
.with_tunnel_generation(node_generation);
|
||||
.with_tunnel_generation(node_generation)
|
||||
.with_settings(settings);
|
||||
let conn = match (security_key.clone(), management_token_credential) {
|
||||
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
||||
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
||||
@@ -416,19 +447,22 @@ async fn run_proxy_reader(
|
||||
let idle_enabled = !idle_timeout.is_zero();
|
||||
let mut oversized_count = 0u32;
|
||||
let mut frames_received: u64 = 0;
|
||||
let mut close_rx = conn.outbound.subscribe_close();
|
||||
let mut heartbeats = JoinSet::new();
|
||||
loop {
|
||||
let msg = if idle_enabled {
|
||||
tokio::select! {
|
||||
msg = ws_rx.next() => msg,
|
||||
_ = tokio::time::sleep(idle_timeout) => {
|
||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
if conn.outbound.is_closing() {
|
||||
break;
|
||||
}
|
||||
while heartbeats.try_join_next().is_some() {}
|
||||
let msg = tokio::select! {
|
||||
biased;
|
||||
_ = close_rx.changed() => break,
|
||||
msg = ws_rx.next() => msg,
|
||||
_ = tokio::time::sleep(idle_timeout), if idle_enabled => {
|
||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
ws_rx.next().await
|
||||
};
|
||||
|
||||
match msg {
|
||||
@@ -463,7 +497,27 @@ async fn run_proxy_reader(
|
||||
continue;
|
||||
}
|
||||
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
let is_heartbeat = protocol::FrameHeader::parse(&data)
|
||||
.is_some_and(|header| header.msg_type == protocol::HEARTBEAT_DATA);
|
||||
if is_heartbeat {
|
||||
if heartbeats.is_empty() {
|
||||
let heartbeat_hub = Arc::clone(&hub);
|
||||
let conn_id = conn.id;
|
||||
heartbeats.spawn(async move {
|
||||
if tokio::time::timeout(
|
||||
Duration::from_secs(10),
|
||||
heartbeat_hub.handle_proxy_frame(conn_id, &mut data),
|
||||
)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
warn!(conn_id, "proxy heartbeat processing timed out");
|
||||
}
|
||||
});
|
||||
}
|
||||
} else {
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(
|
||||
@@ -489,6 +543,56 @@ async fn run_proxy_reader(
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
heartbeats.shutdown().await;
|
||||
}
|
||||
|
||||
async fn read_proxy_settings(
|
||||
ws_tx: &mut futures_util::stream::SplitSink<WebSocket, Message>,
|
||||
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
|
||||
security: Option<&SecureFrameCodec>,
|
||||
protocol_version: u8,
|
||||
) -> Option<protocol::SettingsPayload> {
|
||||
tokio::time::timeout(PROXY_HELLO_TIMEOUT, async {
|
||||
let mut hello_received = security.is_some();
|
||||
for _ in 0..MAX_PREAUTH_PINGS {
|
||||
match ws_rx.next().await? {
|
||||
Ok(Message::Binary(data)) => {
|
||||
if data.len() > 256 * 1024 {
|
||||
return None;
|
||||
}
|
||||
let data = decrypt_message(data, security).ok()?;
|
||||
let frame = Frame::decode(data.into()).ok()?;
|
||||
if frame.stream_id != 0 || frame.flags != 0 {
|
||||
return None;
|
||||
}
|
||||
match frame.msg_type {
|
||||
MsgType::Hello if !hello_received => {
|
||||
let hello =
|
||||
serde_json::from_slice::<HelloPayload>(&frame.payload).ok()?;
|
||||
if hello.protocol_version != protocol_version {
|
||||
return None;
|
||||
}
|
||||
hello_received = true;
|
||||
}
|
||||
MsgType::Settings if hello_received => {
|
||||
let settings =
|
||||
serde_json::from_slice::<protocol::SettingsPayload>(&frame.payload)
|
||||
.ok()?;
|
||||
return settings.is_valid().then_some(settings);
|
||||
}
|
||||
_ => return None,
|
||||
}
|
||||
}
|
||||
Ok(Message::Ping(payload)) => ws_tx.send(Message::Pong(payload)).await.ok()?,
|
||||
Ok(Message::Pong(_)) => {}
|
||||
_ => return None,
|
||||
}
|
||||
}
|
||||
None
|
||||
})
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn encrypt_message(
|
||||
@@ -523,6 +627,158 @@ fn decrypt_message(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[cfg(feature = "testkit")]
|
||||
#[tokio::test]
|
||||
async fn slow_heartbeat_does_not_block_response_frames() {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message as ClientMessage};
|
||||
|
||||
let called = Arc::new(AtomicUsize::new(0));
|
||||
let callback_called = Arc::clone(&called);
|
||||
let control_plane = super::super::control_plane::ControlPlaneClient::local(
|
||||
move |_, _| {
|
||||
callback_called.fetch_add(1, Ordering::SeqCst);
|
||||
Box::pin(std::future::pending())
|
||||
},
|
||||
|_, _, _, _| Box::pin(async { Ok(()) }),
|
||||
);
|
||||
let data = crate::data::GatewayDataState::with_tunnel_management_auth_for_testkit(
|
||||
"heartbeat-test",
|
||||
"heartbeat-generation",
|
||||
"ae-tunnel-harness-management-token",
|
||||
aether_crypto::DEVELOPMENT_ENCRYPTION_KEY,
|
||||
)
|
||||
.unwrap();
|
||||
let state = super::super::AppState::new(
|
||||
control_plane,
|
||||
ConnConfig {
|
||||
ping_interval: Duration::from_secs(60),
|
||||
idle_timeout: Duration::ZERO,
|
||||
outbound_queue_capacity: 128,
|
||||
},
|
||||
16,
|
||||
)
|
||||
.with_data(Arc::new(data));
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let router = super::super::build_router_with_state(state.clone());
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(
|
||||
listener,
|
||||
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
});
|
||||
let mut request = format!("ws://{address}/api/internal/proxy-tunnel")
|
||||
.into_client_request()
|
||||
.unwrap();
|
||||
let headers = request.headers_mut();
|
||||
headers.insert("x-node-id", "heartbeat-test".parse().unwrap());
|
||||
headers.insert(
|
||||
aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER,
|
||||
"heartbeat-generation".parse().unwrap(),
|
||||
);
|
||||
headers.insert(
|
||||
aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER,
|
||||
"3".parse().unwrap(),
|
||||
);
|
||||
headers.insert(
|
||||
"authorization",
|
||||
"Bearer ae-tunnel-harness-management-token".parse().unwrap(),
|
||||
);
|
||||
let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap();
|
||||
let hello = HelloPayload {
|
||||
protocol_version: 3,
|
||||
capabilities: vec![],
|
||||
session_id: None,
|
||||
replica_id: None,
|
||||
};
|
||||
websocket
|
||||
.send(ClientMessage::Binary(protocol::encode_hello(&hello).into()))
|
||||
.await
|
||||
.unwrap();
|
||||
websocket
|
||||
.send(ClientMessage::Binary(
|
||||
protocol::encode_settings(&super::super::hub::local_settings()).into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let ClientMessage::Binary(settings) = websocket.next().await.unwrap().unwrap() else {
|
||||
panic!("expected SETTINGS")
|
||||
};
|
||||
assert_eq!(Frame::decode(settings).unwrap().msg_type, MsgType::Settings);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while !state.hub.has_local_proxy("heartbeat-test") {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let meta: protocol::RequestMeta = serde_json::from_value(serde_json::json!({
|
||||
"method": "GET", "url": "https://example.com", "headers": {}, "stream": true, "timeout": 10
|
||||
})).unwrap();
|
||||
let stream = state
|
||||
.hub
|
||||
.open_local_stream("heartbeat-test", &meta)
|
||||
.await
|
||||
.unwrap();
|
||||
let ClientMessage::Binary(request) = websocket.next().await.unwrap().unwrap() else {
|
||||
panic!("expected request headers")
|
||||
};
|
||||
let stream_id = Frame::decode(request).unwrap().stream_id;
|
||||
let heartbeat = Frame::control(
|
||||
MsgType::HeartbeatData,
|
||||
serde_json::to_vec(&serde_json::json!({"node_id": "heartbeat-test"})).unwrap(),
|
||||
);
|
||||
websocket
|
||||
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while called.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
for _ in 0..8 {
|
||||
websocket
|
||||
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let response = Frame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseHeaders,
|
||||
0,
|
||||
serde_json::to_vec(&serde_json::json!({"status": 200, "headers": []})).unwrap(),
|
||||
);
|
||||
websocket
|
||||
.send(ClientMessage::Binary(response.encode()))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
stream
|
||||
.wait_headers(Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap()
|
||||
.status,
|
||||
200
|
||||
);
|
||||
assert_eq!(called.load(Ordering::SeqCst), 1);
|
||||
state.hub.request_close_all_proxies();
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while state.hub.has_local_proxy("heartbeat-test") {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
server.abort();
|
||||
let _ = server.await;
|
||||
}
|
||||
|
||||
use super::*;
|
||||
|
||||
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||
|
||||
Reference in New Issue
Block a user