Refactor tunnel stability protocol

This commit is contained in:
elky
2026-06-01 01:36:49 +08:00
parent 392353ffff
commit 37413c0211
21 changed files with 1960 additions and 132 deletions
Generated
+1 -1
View File
@@ -480,7 +480,7 @@ dependencies = [
[[package]] [[package]]
name = "aether-tunnel" name = "aether-tunnel"
version = "0.3.15" version = "0.3.16"
dependencies = [ dependencies = [
"aether-contracts", "aether-contracts",
"aether-gateway", "aether-gateway",
File diff suppressed because it is too large Load Diff
@@ -88,8 +88,13 @@ pub(crate) async fn open_direct_relay_stream(
let stream = state let stream = state
.hub .hub
.open_local_stream(node_id, &meta) .open_local_stream(node_id, &meta)
.await
.map_err(|error| format!("connect: {error}"))?; .map_err(|error| format!("connect: {error}"))?;
if let Err(error) = state.hub.push_local_request_body(stream.id, body, true) { if let Err(error) = state
.hub
.push_local_request_body(stream.id, body, true)
.await
{
state.hub.cancel_local_stream(stream.id, &error); state.hub.cancel_local_stream(stream.id, &error);
return Err(format!("connect: {error}")); return Err(format!("connect: {error}"));
} }
@@ -259,7 +264,7 @@ pub async fn relay_request(
continue; continue;
}; };
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) { let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta).await {
Ok(stream) => stream, Ok(stream) => stream,
Err(error) => { Err(error) => {
return release_permit_response( return release_permit_response(
@@ -271,10 +276,10 @@ pub async fn relay_request(
if envelope_buf.len() > body_offset { if envelope_buf.len() > body_offset {
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]); let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
if let Err(error) = if let Err(error) = state
state .hub
.hub .push_local_request_body(opened_stream.id, first_body_chunk, false)
.push_local_request_body(opened_stream.id, first_body_chunk, false) .await
{ {
state.hub.cancel_local_stream(opened_stream.id, &error); state.hub.cancel_local_stream(opened_stream.id, &error);
return release_permit_response( return release_permit_response(
@@ -296,6 +301,7 @@ pub async fn relay_request(
if let Err(error) = state if let Err(error) = state
.hub .hub
.push_local_request_body(active_stream.id, chunk, false) .push_local_request_body(active_stream.id, chunk, false)
.await
{ {
state.hub.cancel_local_stream(active_stream.id, &error); state.hub.cancel_local_stream(active_stream.id, &error);
return release_permit_response( return release_permit_response(
@@ -322,6 +328,7 @@ pub async fn relay_request(
if let Err(error) = state if let Err(error) = state
.hub .hub
.push_local_request_body(stream.id, Bytes::new(), true) .push_local_request_body(stream.id, Bytes::new(), true)
.await
{ {
state.hub.cancel_local_stream(stream.id, &error); state.hub.cancel_local_stream(stream.id, &error);
return release_permit_response( return release_permit_response(
@@ -376,6 +376,7 @@ fn resolve_proxy_protocol_version(headers: &HeaderMap) -> u8 {
.and_then(|value| value.to_str().ok()) .and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u8>().ok()) .and_then(|value| value.parse::<u8>().ok())
.filter(|value| *value >= 1) .filter(|value| *value >= 1)
.map(|value| value.min(aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION))
.unwrap_or(1) .unwrap_or(1)
} }
@@ -1,10 +1,14 @@
use bytes::Bytes; use bytes::Bytes;
pub use aether_contracts::tunnel::{ pub use aether_contracts::tunnel::{
decode_payload, encode_frame, encode_goaway, encode_ping, encode_pong, encode_stream_error, decode_payload, encode_connection_close, encode_frame, encode_goaway, encode_goaway_v3,
frame_payload_by_header, FrameHeader, RequestMeta, ResponseMeta, FLAG_END_STREAM, encode_hello, encode_load_report, encode_ping, encode_pong, encode_reset_stream,
FLAG_GZIP_COMPRESSED, GOAWAY, HEADER_SIZE, HEARTBEAT_ACK, HEARTBEAT_DATA, PING, PONG, encode_settings, encode_stream_error, encode_window_update, frame_payload_by_header,
REQUEST_BODY, REQUEST_HEADERS, RESPONSE_BODY, RESPONSE_HEADERS, STREAM_END, STREAM_ERROR, ConnectionClosePayload, FrameHeader, GoAwayPayload, HelloPayload, LoadReportPayload,
RequestMeta, ResetStreamPayload, ResponseMeta, SettingsPayload, WindowUpdatePayload,
CONNECTION_CLOSE, FLAG_END_STREAM, FLAG_GZIP_COMPRESSED, GOAWAY, HEADER_SIZE, HEARTBEAT_ACK,
HEARTBEAT_DATA, HELLO, LOAD_REPORT, PING, PONG, REQUEST_BODY, REQUEST_HEADERS, RESET_STREAM,
RESPONSE_BODY, RESPONSE_HEADERS, SETTINGS, STREAM_END, STREAM_ERROR, WINDOW_UPDATE,
}; };
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> { pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
@@ -182,7 +182,9 @@ pub async fn handle_proxy_connection(
loop { loop {
tokio::time::sleep(ping_interval).await; tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping(); let ping = protocol::encode_ping();
let status = ping_conn.send(Message::Binary(ping.into())); let status = ping_conn
.send_wait(Message::Binary(ping.into()), Duration::from_millis(250))
.await;
if !matches!(status, SendStatus::Queued) { if !matches!(status, SendStatus::Queued) {
let snapshot = ping_conn.outbound.snapshot(); let snapshot = ping_conn.outbound.snapshot();
match status { match status {
+3 -2
View File
@@ -535,12 +535,13 @@ impl EmbeddedTunnelState {
http1_only: false, http1_only: false,
transport_profile: None, transport_profile: None,
}; };
let stream = self.inner.hub.open_local_stream(node_id, &meta)?; let stream = self.inner.hub.open_local_stream(node_id, &meta).await?;
let stream_id = stream.id; let stream_id = stream.id;
let result = async { let result = async {
self.inner self.inner
.hub .hub
.push_local_request_body(stream_id, Bytes::new(), true)?; .push_local_request_body(stream_id, Bytes::new(), true)
.await?;
let response = stream let response = stream
.wait_headers(Duration::from_secs(timeout_secs)) .wait_headers(Duration::from_secs(timeout_secs))
.await?; .await?;
+1 -1
View File
@@ -10,7 +10,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
[[package]] [[package]]
name = "aether-tunnel" name = "aether-tunnel"
version = "0.3.12" version = "0.3.16"
dependencies = [ dependencies = [
"anyhow", "anyhow",
"arc-swap", "arc-swap",
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "aether-tunnel" name = "aether-tunnel"
version = "0.3.15" version = "0.3.16"
edition = "2021" edition = "2021"
description = "Tunnel agent for Aether" description = "Tunnel agent for Aether"
+4 -1
View File
@@ -128,6 +128,9 @@ sudo aether-tunnel uninstall
| `--tunnel-connections` | `AETHER_TUNNEL_CONNECTIONS` | 自动(硬件估算) | 最小连接池大小;显式设置后默认固定为该值 | | `--tunnel-connections` | `AETHER_TUNNEL_CONNECTIONS` | 自动(硬件估算) | 最小连接池大小;显式设置后默认固定为该值 |
| `--tunnel-connections-max` | `AETHER_TUNNEL_CONNECTIONS_MAX` | 自动(硬件估算) | 连接池自动扩容上限;大于 `tunnel_connections` 时启用 autoscale | | `--tunnel-connections-max` | `AETHER_TUNNEL_CONNECTIONS_MAX` | 自动(硬件估算) | 连接池自动扩容上限;大于 `tunnel_connections` 时启用 autoscale |
| `--tunnel-max-streams` | `AETHER_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 | | `--tunnel-max-streams` | `AETHER_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
| `--tunnel-profile` | `AETHER_TUNNEL_PROFILE` | `standard` | 自动连接池档位:`lite=2`、`standard=4`、`throughput=8+autoscale` |
| `--tunnel-stream-initial-window-bytes` | `AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES` | `4194304` | v3 stream 初始流控窗口 |
| `--tunnel-drain-deadline-ms` | `AETHER_TUNNEL_DRAIN_DEADLINE_MS` | `30000` | v3 GOAWAY/drain 优雅退出期限 |
| `--tunnel-ping-interval-ms` | `AETHER_TUNNEL_PING_INTERVAL_MS` | `10000` | WebSocket ping 周期(毫秒) | | `--tunnel-ping-interval-ms` | `AETHER_TUNNEL_PING_INTERVAL_MS` | `10000` | WebSocket ping 周期(毫秒) |
| `--tunnel-connect-timeout-ms` | `AETHER_TUNNEL_CONNECT_TIMEOUT_MS` | `3000` | tunnel 建连超时(毫秒) | | `--tunnel-connect-timeout-ms` | `AETHER_TUNNEL_CONNECT_TIMEOUT_MS` | `3000` | tunnel 建连超时(毫秒) |
| `--tunnel-ipv4-only` | `AETHER_TUNNEL_IPV4_ONLY` | `false` | 仅使用 IPv4 地址建立直连 WebSocket tunnel;配置 `aether_outbound_proxy_url` 时仅限制代理端点解析 | | `--tunnel-ipv4-only` | `AETHER_TUNNEL_IPV4_ONLY` | `false` | 仅使用 IPv4 地址建立直连 WebSocket tunnel;配置 `aether_outbound_proxy_url` 时仅限制代理端点解析 |
@@ -142,7 +145,7 @@ sudo aether-tunnel uninstall
| `--tunnel-reconnect-base-ms` | `AETHER_TUNNEL_RECONNECT_BASE_MS` | `50` | 指数退避基础延迟(毫秒) | | `--tunnel-reconnect-base-ms` | `AETHER_TUNNEL_RECONNECT_BASE_MS` | `50` | 指数退避基础延迟(毫秒) |
| `--tunnel-reconnect-max-ms` | `AETHER_TUNNEL_RECONNECT_MAX_MS` | `250` | 指数退避上限(毫秒) | | `--tunnel-reconnect-max-ms` | `AETHER_TUNNEL_RECONNECT_MAX_MS` | `250` | 指数退避上限(毫秒) |
省略 `tunnel_connections` 时,tunnel 会按设备能力自动计算一个基线值和偏单机上限的扩容上限:默认至少保留 2 条常驻 tunnel,并会更早触发扩容;如果显式设置了 `tunnel_connections` 但没有设置 `tunnel_connections_max`,则保持固定连接池,不自动扩缩。 省略 `tunnel_connections` 时,tunnel 会按 `tunnel_profile` 和设备能力自动计算一个基线值和扩容上限:`standard` 默认至少保留 4 条常驻 tunnel;如果显式设置了 `tunnel_connections` 但没有设置 `tunnel_connections_max`,则保持固定连接池,不自动扩缩。
`tunnel_ipv4_only` / `tunnel_ipv6_only` 只能二选一。它们只改变 WebSocket tunnel 回连的 TCP 地址选择:直连 Aether 时过滤 Aether 域名的 DNS 结果;配置 `aether_outbound_proxy_url` 时过滤代理服务器端点的 DNS 结果,Host/SNI 仍使用原始 WebSocket URL。该选项不会影响 provider 上游请求;如需限制 provider 上游流量,请在 `upstream_proxy_url` 或系统网络层处理。对于 Cloudflare 等边缘 IP 会变化的域名,优先使用该选项而不是固定 `/etc/hosts`。 `tunnel_ipv4_only` / `tunnel_ipv6_only` 只能二选一。它们只改变 WebSocket tunnel 回连的 TCP 地址选择:直连 Aether 时过滤 Aether 域名的 DNS 结果;配置 `aether_outbound_proxy_url` 时过滤代理服务器端点的 DNS 结果,Host/SNI 仍使用原始 WebSocket URL。该选项不会影响 provider 上游请求;如需限制 provider 上游流量,请在 `upstream_proxy_url` 或系统网络层处理。对于 Cloudflare 等边缘 IP 会变化的域名,优先使用该选项而不是固定 `/etc/hosts`。
+14 -1
View File
@@ -171,6 +171,9 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
tunnel_connections_initial = tunnel_pool_policy.min_connections, tunnel_connections_initial = tunnel_pool_policy.min_connections,
tunnel_connections_max = tunnel_pool_policy.max_connections, tunnel_connections_max = tunnel_pool_policy.max_connections,
tunnel_max_streams = tunnel_pool_policy.max_streams_per_tunnel, tunnel_max_streams = tunnel_pool_policy.max_streams_per_tunnel,
tunnel_profile = %config.tunnel_profile,
tunnel_stream_initial_window_bytes = config.tunnel_stream_initial_window_bytes,
tunnel_drain_deadline_ms = config.tunnel_drain_deadline_ms,
scale_check_interval_ms = tunnel_pool_policy.scale_check_interval.as_millis(), scale_check_interval_ms = tunnel_pool_policy.scale_check_interval.as_millis(),
scale_up_threshold_percent = tunnel_pool_policy.scale_up_threshold_percent, scale_up_threshold_percent = tunnel_pool_policy.scale_up_threshold_percent,
scale_down_threshold_percent = tunnel_pool_policy.scale_down_threshold_percent, scale_down_threshold_percent = tunnel_pool_policy.scale_down_threshold_percent,
@@ -490,6 +493,9 @@ async fn diagnostics_stats(
"max_in_flight_streams": diagnostics.state.config.max_in_flight_streams, "max_in_flight_streams": diagnostics.state.config.max_in_flight_streams,
"distributed_stream_limit": diagnostics.state.config.distributed_stream_limit, "distributed_stream_limit": diagnostics.state.config.distributed_stream_limit,
"tunnel_max_streams": diagnostics.state.config.tunnel_max_streams, "tunnel_max_streams": diagnostics.state.config.tunnel_max_streams,
"tunnel_profile": diagnostics.state.config.tunnel_profile.to_string(),
"tunnel_stream_initial_window_bytes": diagnostics.state.config.tunnel_stream_initial_window_bytes,
"tunnel_drain_deadline_ms": diagnostics.state.config.tunnel_drain_deadline_ms,
"tunnel_connections": diagnostics.state.config.tunnel_connections, "tunnel_connections": diagnostics.state.config.tunnel_connections,
"tunnel_connections_max": diagnostics.state.config.tunnel_connections_max, "tunnel_connections_max": diagnostics.state.config.tunnel_connections_max,
"diagnostics_bind": diagnostics.state.config.diagnostics_bind.map(|addr| addr.to_string()), "diagnostics_bind": diagnostics.state.config.diagnostics_bind.map(|addr| addr.to_string()),
@@ -1131,7 +1137,10 @@ mod tests {
.await .await
.expect("stats response should parse"); .expect("stats response should parse");
assert_eq!(stats["status"], "ok"); assert_eq!(stats["status"], "ok");
assert_eq!(stats["protocol_version"], 2); assert_eq!(
stats["protocol_version"],
aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION
);
assert_eq!(stats["servers"][0]["node_id"], "node-diagnostics"); assert_eq!(stats["servers"][0]["node_id"], "node-diagnostics");
handle.abort(); handle.abort();
@@ -1400,6 +1409,10 @@ mod tests {
tunnel_reconnect_max_ms: 250, tunnel_reconnect_max_ms: 250,
tunnel_ping_interval_ms: 1_000, tunnel_ping_interval_ms: 1_000,
tunnel_max_streams: Some(8), tunnel_max_streams: Some(8),
tunnel_profile: crate::config::TunnelProfileArg::Lite,
tunnel_stream_initial_window_bytes:
crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES,
tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS,
tunnel_connect_timeout_ms: 2_000, tunnel_connect_timeout_ms: 2_000,
tunnel_ipv4_only: false, tunnel_ipv4_only: false,
tunnel_ipv6_only: false, tunnel_ipv6_only: false,
+89 -10
View File
@@ -65,6 +65,8 @@ pub const DEFAULT_TUNNEL_SCALE_CHECK_INTERVAL_MS: u64 = 1_000;
pub const DEFAULT_TUNNEL_SCALE_UP_THRESHOLD_PERCENT: u32 = 50; pub const DEFAULT_TUNNEL_SCALE_UP_THRESHOLD_PERCENT: u32 = 50;
pub const DEFAULT_TUNNEL_SCALE_DOWN_THRESHOLD_PERCENT: u32 = 35; pub const DEFAULT_TUNNEL_SCALE_DOWN_THRESHOLD_PERCENT: u32 = 35;
pub const DEFAULT_TUNNEL_SCALE_DOWN_GRACE_SECS: u64 = 15; pub const DEFAULT_TUNNEL_SCALE_DOWN_GRACE_SECS: u64 = 15;
pub const DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES: u32 = 4 * 1024 * 1024;
pub const DEFAULT_TUNNEL_DRAIN_DEADLINE_MS: u64 = 30_000;
pub const DEFAULT_UPSTREAM_CLIENT_POOL_CAPACITY: usize = 256; pub const DEFAULT_UPSTREAM_CLIENT_POOL_CAPACITY: usize = 256;
const AUTO_TUNNEL_CONNECTIONS_REDUNDANT_FLOOR: u64 = 2; const AUTO_TUNNEL_CONNECTIONS_REDUNDANT_FLOOR: u64 = 2;
const AUTO_TUNNEL_CONNECTIONS_BASE_CAP: u64 = 4; const AUTO_TUNNEL_CONNECTIONS_BASE_CAP: u64 = 4;
@@ -76,6 +78,9 @@ const AUTO_TUNNEL_CONNECTIONS_MAX_CAP: u64 = 32;
const TUNNEL_PING_INTERVAL_MS_ENV: &str = "AETHER_TUNNEL_PING_INTERVAL_MS"; const TUNNEL_PING_INTERVAL_MS_ENV: &str = "AETHER_TUNNEL_PING_INTERVAL_MS";
const TUNNEL_CONNECT_TIMEOUT_MS_ENV: &str = "AETHER_TUNNEL_CONNECT_TIMEOUT_MS"; const TUNNEL_CONNECT_TIMEOUT_MS_ENV: &str = "AETHER_TUNNEL_CONNECT_TIMEOUT_MS";
const TUNNEL_STALE_TIMEOUT_MS_ENV: &str = "AETHER_TUNNEL_STALE_TIMEOUT_MS"; const TUNNEL_STALE_TIMEOUT_MS_ENV: &str = "AETHER_TUNNEL_STALE_TIMEOUT_MS";
const TUNNEL_PROFILE_ENV: &str = "AETHER_TUNNEL_PROFILE";
const TUNNEL_STREAM_INITIAL_WINDOW_BYTES_ENV: &str = "AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES";
const TUNNEL_DRAIN_DEADLINE_MS_ENV: &str = "AETHER_TUNNEL_DRAIN_DEADLINE_MS";
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TunnelPoolSizing { pub struct TunnelPoolSizing {
@@ -249,6 +254,24 @@ impl From<TunnelLogRotationArg> for LogRotation {
} }
} }
#[derive(clap::ValueEnum, Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum TunnelProfileArg {
Lite,
Standard,
Throughput,
}
impl fmt::Display for TunnelProfileArg {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
TunnelProfileArg::Lite => "lite",
TunnelProfileArg::Standard => "standard",
TunnelProfileArg::Throughput => "throughput",
})
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum TunnelSecurity { pub enum TunnelSecurity {
@@ -642,6 +665,31 @@ pub struct Config {
#[arg(long, env = "AETHER_TUNNEL_MAX_STREAMS")] #[arg(long, env = "AETHER_TUNNEL_MAX_STREAMS")]
pub tunnel_max_streams: Option<u32>, pub tunnel_max_streams: Option<u32>,
/// Tunnel connection pool profile used when connection counts are not explicit.
#[arg(
long,
env = TUNNEL_PROFILE_ENV,
value_enum,
default_value_t = TunnelProfileArg::Standard
)]
pub tunnel_profile: TunnelProfileArg,
/// Initial per-stream flow-control window advertised by this tunnel.
#[arg(
long,
env = TUNNEL_STREAM_INITIAL_WINDOW_BYTES_ENV,
default_value_t = DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES
)]
pub tunnel_stream_initial_window_bytes: u32,
/// Deadline for graceful tunnel drain after GOAWAY.
#[arg(
long,
env = TUNNEL_DRAIN_DEADLINE_MS_ENV,
default_value_t = DEFAULT_TUNNEL_DRAIN_DEADLINE_MS
)]
pub tunnel_drain_deadline_ms: u64,
/// WebSocket tunnel TCP connect timeout in milliseconds /// WebSocket tunnel TCP connect timeout in milliseconds
#[arg( #[arg(
long, long,
@@ -787,6 +835,12 @@ impl Config {
if matches!(self.tunnel_connections_max, Some(0)) { if matches!(self.tunnel_connections_max, Some(0)) {
anyhow::bail!("tunnel_connections_max must be > 0"); anyhow::bail!("tunnel_connections_max must be > 0");
} }
if self.tunnel_stream_initial_window_bytes == 0 {
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
}
if self.tunnel_drain_deadline_ms == 0 {
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
}
if let (Some(min_connections), Some(max_connections)) = if let (Some(min_connections), Some(max_connections)) =
(self.tunnel_connections, self.tunnel_connections_max) (self.tunnel_connections, self.tunnel_connections_max)
{ {
@@ -908,14 +962,20 @@ impl Config {
.and_then(|limit| u64::try_from(limit).ok()) .and_then(|limit| u64::try_from(limit).ok())
.unwrap_or(hw_info.estimated_max_concurrency) .unwrap_or(hw_info.estimated_max_concurrency)
.max(per_tunnel_capacity); .max(per_tunnel_capacity);
let (profile_initial_floor, profile_initial_cap, profile_max_floor) =
match self.tunnel_profile {
TunnelProfileArg::Lite => (2, 2, 2),
TunnelProfileArg::Standard => (4, 4, 4),
TunnelProfileArg::Throughput => (8, 8, 8),
};
let cpu_soft_cap = u64::from(hw_info.cpu_cores.max(1)) let cpu_soft_cap = u64::from(hw_info.cpu_cores.max(1))
.saturating_mul(AUTO_TUNNEL_CONNECTIONS_PER_CPU_CAP) .saturating_mul(AUTO_TUNNEL_CONNECTIONS_PER_CPU_CAP)
.clamp( .clamp(profile_max_floor, AUTO_TUNNEL_CONNECTIONS_MAX_CAP);
AUTO_TUNNEL_CONNECTIONS_BASE_CAP, let auto_initial_floor = AUTO_TUNNEL_CONNECTIONS_REDUNDANT_FLOOR
AUTO_TUNNEL_CONNECTIONS_MAX_CAP, .max(profile_initial_floor)
); .min(cpu_soft_cap);
let auto_initial_floor = AUTO_TUNNEL_CONNECTIONS_REDUNDANT_FLOOR.min(cpu_soft_cap);
let auto_initial_cap = AUTO_TUNNEL_CONNECTIONS_BASE_CAP let auto_initial_cap = AUTO_TUNNEL_CONNECTIONS_BASE_CAP
.max(profile_initial_cap)
.min(cpu_soft_cap) .min(cpu_soft_cap)
.max(auto_initial_floor); .max(auto_initial_floor);
@@ -926,7 +986,11 @@ impl Config {
100, 100,
) )
.max(1); .max(1);
let auto_max_floor = auto_initial.max(AUTO_TUNNEL_CONNECTIONS_BASE_CAP.min(cpu_soft_cap)); let auto_max_floor = auto_initial.max(
AUTO_TUNNEL_CONNECTIONS_BASE_CAP
.max(profile_max_floor)
.min(cpu_soft_cap),
);
let auto_max = let auto_max =
div_ceil_u64(estimated, high_water_per_tunnel).clamp(auto_max_floor, cpu_soft_cap); div_ceil_u64(estimated, high_water_per_tunnel).clamp(auto_max_floor, cpu_soft_cap);
@@ -1092,6 +1156,12 @@ pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_max_streams: Option<u32>, pub tunnel_max_streams: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_profile: Option<TunnelProfileArg>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_stream_initial_window_bytes: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_drain_deadline_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_connect_timeout_ms: Option<u64>, pub tunnel_connect_timeout_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub tunnel_ipv4_only: Option<bool>, pub tunnel_ipv4_only: Option<bool>,
@@ -1291,6 +1361,15 @@ impl ConfigFile {
); );
set!(TUNNEL_PING_INTERVAL_MS_ENV, self.tunnel_ping_interval_ms); set!(TUNNEL_PING_INTERVAL_MS_ENV, self.tunnel_ping_interval_ms);
set!("AETHER_TUNNEL_MAX_STREAMS", self.tunnel_max_streams); set!("AETHER_TUNNEL_MAX_STREAMS", self.tunnel_max_streams);
set!(
TUNNEL_PROFILE_ENV,
self.tunnel_profile.map(|value| value.to_string())
);
set!(
TUNNEL_STREAM_INITIAL_WINDOW_BYTES_ENV,
self.tunnel_stream_initial_window_bytes
);
set!(TUNNEL_DRAIN_DEADLINE_MS_ENV, self.tunnel_drain_deadline_ms);
set!( set!(
TUNNEL_CONNECT_TIMEOUT_MS_ENV, TUNNEL_CONNECT_TIMEOUT_MS_ENV,
self.tunnel_connect_timeout_ms self.tunnel_connect_timeout_ms
@@ -2056,7 +2135,7 @@ node_name = "tunnel-test"
let sizing = config let sizing = config
.resolve_tunnel_pool_sizing(&hw) .resolve_tunnel_pool_sizing(&hw)
.expect("sizing should resolve"); .expect("sizing should resolve");
assert_eq!(sizing.initial_connections, 3); assert_eq!(sizing.initial_connections, 4);
assert_eq!(sizing.max_connections, 32); assert_eq!(sizing.max_connections, 32);
} }
@@ -2084,7 +2163,7 @@ node_name = "tunnel-test"
let sizing = config let sizing = config
.resolve_tunnel_pool_sizing(&hw) .resolve_tunnel_pool_sizing(&hw)
.expect("sizing should resolve"); .expect("sizing should resolve");
assert_eq!(sizing.initial_connections, 2); assert_eq!(sizing.initial_connections, 4);
assert_eq!(sizing.max_connections, 4); assert_eq!(sizing.max_connections, 4);
} }
@@ -2112,7 +2191,7 @@ node_name = "tunnel-test"
let sizing = config let sizing = config
.resolve_tunnel_pool_sizing(&hw) .resolve_tunnel_pool_sizing(&hw)
.expect("sizing should resolve"); .expect("sizing should resolve");
assert_eq!(sizing.initial_connections, 2); assert_eq!(sizing.initial_connections, 4);
assert_eq!(sizing.max_connections, 4); assert_eq!(sizing.max_connections, 4);
} }
@@ -2142,7 +2221,7 @@ node_name = "tunnel-test"
let sizing = config let sizing = config
.resolve_tunnel_pool_sizing(&hw) .resolve_tunnel_pool_sizing(&hw)
.expect("sizing should resolve"); .expect("sizing should resolve");
assert_eq!(sizing.initial_connections, 2); assert_eq!(sizing.initial_connections, 4);
assert_eq!(sizing.max_connections, 4); assert_eq!(sizing.max_connections, 4);
} }
+53 -3
View File
@@ -18,7 +18,8 @@ use crate::egress_proxy::{
}; };
use crate::state::{AppState, ServerContext}; use crate::state::{AppState, ServerContext};
use aether_contracts::tunnel::{ use aether_contracts::tunnel::{
CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_NODE_NAME_B64_HEADER, TUNNEL_PROTOCOL_VERSION_HEADER, HelloPayload, SettingsPayload, CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_NODE_NAME_B64_HEADER,
TUNNEL_PROTOCOL_VERSION_HEADER,
}; };
use aether_contracts::tunnel_security::{ use aether_contracts::tunnel_security::{
SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED, SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED,
@@ -186,7 +187,13 @@ pub async fn connect_and_run(
Some(Arc::clone(&server.tunnel_metrics)), Some(Arc::clone(&server.tunnel_metrics)),
security.clone(), security.clone(),
); );
let drain_signal = spawn_drain_signal(conn_idx, frame_tx.clone(), drain.clone()); send_protocol_v3_hello(&frame_tx, &security_session, state).await;
let drain_signal = spawn_drain_signal(
conn_idx,
frame_tx.clone(),
drain.clone(),
state.config.tunnel_drain_deadline_ms,
);
// Spawn heartbeat task (only for primary connection to avoid // Spawn heartbeat task (only for primary connection to avoid
// resetting shared atomic metrics via swap(0)) // resetting shared atomic metrics via swap(0))
@@ -298,10 +305,48 @@ pub async fn connect_and_run(
outcome outcome
} }
async fn send_protocol_v3_hello(
frame_tx: &writer::FrameSender,
security_session: &str,
state: &Arc<AppState>,
) {
let hello = super::protocol::Frame::control(
super::protocol::MsgType::Hello,
serde_json::to_vec(&HelloPayload {
protocol_version: CURRENT_TUNNEL_PROTOCOL_VERSION,
capabilities: vec![
"flow-control".to_string(),
"reset-stream".to_string(),
"graceful-drain".to_string(),
"load-report".to_string(),
],
session_id: Some(security_session.to_string()),
replica_id: None,
})
.expect("hello payload should serialize"),
);
let settings = super::protocol::Frame::control(
super::protocol::MsgType::Settings,
serde_json::to_vec(&SettingsPayload {
initial_stream_window_bytes: state.config.tunnel_stream_initial_window_bytes,
min_window_update_bytes: state
.config
.tunnel_stream_initial_window_bytes
.saturating_div(4)
.max(1),
drain_deadline_ms: state.config.tunnel_drain_deadline_ms,
})
.expect("settings payload should serialize"),
);
let _ = tokio::time::timeout(Duration::from_millis(250), frame_tx.send(hello)).await;
let _ = tokio::time::timeout(Duration::from_millis(250), frame_tx.send(settings)).await;
}
fn spawn_drain_signal( fn spawn_drain_signal(
conn_idx: usize, conn_idx: usize,
frame_tx: writer::FrameSender, frame_tx: writer::FrameSender,
mut drain: watch::Receiver<bool>, mut drain: watch::Receiver<bool>,
drain_deadline_ms: u64,
) -> tokio::task::JoinHandle<()> { ) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move { tokio::spawn(async move {
if !*drain.borrow() { if !*drain.borrow() {
@@ -320,7 +365,12 @@ fn spawn_drain_signal(
Duration::from_millis(250), Duration::from_millis(250),
frame_tx.send(super::protocol::Frame::control( frame_tx.send(super::protocol::Frame::control(
super::protocol::MsgType::GoAway, super::protocol::MsgType::GoAway,
bytes::Bytes::new(), serde_json::to_vec(&aether_contracts::tunnel::GoAwayPayload {
last_accepted_stream_id: u32::MAX,
drain_deadline_ms,
reason: "tunnel drain requested".to_string(),
})
.expect("goaway payload should serialize"),
)), )),
) )
.await .await
+99 -19
View File
@@ -17,6 +17,7 @@ use crate::state::{AppState, ServerContext};
use super::heartbeat::HeartbeatHandle; use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta}; use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::stream_handler; use super::stream_handler;
use super::stream_handler::StreamSendWindow;
use super::writer::FrameSender; use super::writer::FrameSender;
use aether_contracts::tunnel_security::SecureFrameCodec; use aether_contracts::tunnel_security::SecureFrameCodec;
@@ -27,6 +28,12 @@ enum StreamDispatchStatus {
TimedOut, TimedOut,
} }
#[derive(Clone)]
struct StreamDispatchTarget {
body_tx: mpsc::Sender<Frame>,
response_window: Arc<StreamSendWindow>,
}
/// Run the dispatcher loop, reading from the WebSocket stream. /// Run the dispatcher loop, reading from the WebSocket stream.
#[allow(dead_code)] #[allow(dead_code)]
pub async fn run<S>( pub async fn run<S>(
@@ -61,10 +68,11 @@ where
+ Send + Send
+ 'static, + 'static,
{ {
// Active streams: stream_id -> body sender // Active streams: stream_id -> body sender + response flow-control window.
let mut streams: HashMap<u32, mpsc::Sender<Frame>> = HashMap::new(); let mut streams: HashMap<u32, StreamDispatchTarget> = HashMap::new();
// Track spawned stream handlers so we can wait for them on shutdown // Track spawned stream handlers so we can wait for them on shutdown
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new(); let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize; let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
let mut frames_since_cleanup: u32 = 0; let mut frames_since_cleanup: u32 = 0;
let stale_timeout = state let stale_timeout = state
@@ -99,6 +107,16 @@ where
} }
continue; continue;
} }
finished = handler_finished_rx.recv() => {
if let Some(stream_id) = finished {
streams.remove(&stream_id);
if draining && streams.is_empty() {
info!("tunnel drained after stream handler completion");
break None;
}
}
continue;
}
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => { _ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
warn!( warn!(
stale_ms = stale_timeout.as_millis(), stale_ms = stale_timeout.as_millis(),
@@ -239,14 +257,22 @@ where
// Create body channel and spawn handler // Create body channel and spawn handler
let (body_tx, body_rx) = mpsc::channel::<Frame>(64); let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
let response_window = Arc::new(StreamSendWindow::new(
state.config.tunnel_stream_initial_window_bytes,
));
streams.insert(
frame.stream_id,
StreamDispatchTarget {
body_tx,
response_window: Arc::clone(&response_window),
},
);
let request_headers_end_stream = frame.is_end_stream(); let request_headers_end_stream = frame.is_end_stream();
if !request_headers_end_stream {
streams.insert(frame.stream_id, body_tx);
}
let state_clone = Arc::clone(&state); let state_clone = Arc::clone(&state);
let server_clone = Arc::clone(&server); let server_clone = Arc::clone(&server);
let tx_clone = frame_tx.clone(); let tx_clone = frame_tx.clone();
let finished_tx = handler_finished_tx.clone();
let sid = frame.stream_id; let sid = frame.stream_id;
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
stream_handler::handle_stream( stream_handler::handle_stream(
@@ -256,20 +282,33 @@ where
meta, meta,
body_rx, body_rx,
tx_clone, tx_clone,
response_window,
) )
.await; .await;
let _ = finished_tx.send(sid);
}); });
handler_handles.push(handle); handler_handles.push(handle);
if request_headers_end_stream {
if let Some(target) = streams.get(&sid) {
let _ = target.body_tx.try_send(Frame::new(
sid,
MsgType::StreamEnd,
0,
Bytes::new(),
));
}
}
debug!(stream_id = frame.stream_id, "new stream started"); debug!(stream_id = frame.stream_id, "new stream started");
} }
MsgType::RequestBody => { MsgType::RequestBody => {
if let Some(tx) = streams.get(&frame.stream_id).cloned() { if let Some(target) = streams.get(&frame.stream_id).cloned() {
let is_end = frame.is_end_stream(); let is_end = frame.is_end_stream();
let sid = frame.stream_id; let sid = frame.stream_id;
let dispatch = dispatch_stream_frame(&tx, frame).await; let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
if is_end || dispatch != StreamDispatchStatus::Delivered { if dispatch != StreamDispatchStatus::Delivered {
streams.remove(&sid); streams.remove(&sid);
if dispatch == StreamDispatchStatus::TimedOut { if dispatch == StreamDispatchStatus::TimedOut {
server.tunnel_metrics.record_error( server.tunnel_metrics.record_error(
@@ -282,7 +321,7 @@ where
"tunnel request body dispatch stalled", "tunnel request body dispatch stalled",
); );
} }
if draining && streams.is_empty() { if is_end && draining && streams.is_empty() {
info!("tunnel drained after request body completion"); info!("tunnel drained after request body completion");
break None; break None;
} }
@@ -290,10 +329,10 @@ where
} }
} }
MsgType::StreamEnd | MsgType::StreamError => { MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
// Client-side cancellation or end // Client-side cancellation or end
if let Some(tx) = streams.remove(&frame.stream_id) { if let Some(target) = streams.remove(&frame.stream_id) {
let _ = dispatch_stream_frame(&tx, frame).await; let _ = dispatch_stream_frame(&target.body_tx, frame).await;
if draining && streams.is_empty() { if draining && streams.is_empty() {
info!("tunnel drained after stream termination"); info!("tunnel drained after stream termination");
break None; break None;
@@ -320,6 +359,35 @@ where
break None; break None;
} }
MsgType::WindowUpdate => {
if let Ok(payload) = serde_json::from_slice::<
aether_contracts::tunnel::WindowUpdatePayload,
>(&frame.payload)
{
if let Some(target) = streams.get(&frame.stream_id) {
target.response_window.add_credit(payload.delta_bytes);
}
}
debug!(
msg_type = ?frame.msg_type,
stream_id = frame.stream_id,
"received tunnel protocol v3 WINDOW_UPDATE frame"
);
}
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
debug!(
msg_type = ?frame.msg_type,
stream_id = frame.stream_id,
"received tunnel protocol v3 control frame"
);
}
MsgType::ConnectionClose => {
info!("received CONNECTION_CLOSE");
break None;
}
_ => { _ => {
debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type"); debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type");
} }
@@ -329,10 +397,6 @@ where
// Trigger every 64 frames OR when the count exceeds max_streams. // Trigger every 64 frames OR when the count exceeds max_streams.
frames_since_cleanup += 1; frames_since_cleanup += 1;
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams { if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
let closed_streams = prune_closed_stream_senders(&mut streams);
if closed_streams > 0 {
debug!(closed_streams, "removed closed request body stream senders");
}
handler_handles.retain(|h| !h.is_finished()); handler_handles.retain(|h| !h.is_finished());
frames_since_cleanup = 0; frames_since_cleanup = 0;
if draining && streams.is_empty() { if draining && streams.is_empty() {
@@ -408,9 +472,10 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
} }
} }
fn prune_closed_stream_senders(streams: &mut HashMap<u32, mpsc::Sender<Frame>>) -> usize { #[cfg(test)]
fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>) -> usize {
let before = streams.len(); let before = streams.len();
streams.retain(|_, tx| !tx.is_closed()); streams.retain(|_, target| !target.body_tx.is_closed());
before.saturating_sub(streams.len()) before.saturating_sub(streams.len())
} }
@@ -493,7 +558,22 @@ mod tests {
let (closed_tx, closed_rx) = mpsc::channel::<Frame>(1); let (closed_tx, closed_rx) = mpsc::channel::<Frame>(1);
let (open_tx, _open_rx) = mpsc::channel::<Frame>(1); let (open_tx, _open_rx) = mpsc::channel::<Frame>(1);
drop(closed_rx); drop(closed_rx);
let mut streams = HashMap::from([(7, closed_tx), (9, open_tx)]); let mut streams = HashMap::from([
(
7,
StreamDispatchTarget {
body_tx: closed_tx,
response_window: Arc::new(StreamSendWindow::new(1024)),
},
),
(
9,
StreamDispatchTarget {
body_tx: open_tx,
response_window: Arc::new(StreamSendWindow::new(1024)),
},
),
]);
let removed = prune_closed_stream_senders(&mut streams); let removed = prune_closed_stream_senders(&mut streams);
+4
View File
@@ -548,6 +548,10 @@ mod tests {
tunnel_reconnect_max_ms: 250, tunnel_reconnect_max_ms: 250,
tunnel_ping_interval_ms: 1_000, tunnel_ping_interval_ms: 1_000,
tunnel_max_streams: Some(8), tunnel_max_streams: Some(8),
tunnel_profile: crate::config::TunnelProfileArg::Lite,
tunnel_stream_initial_window_bytes:
crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES,
tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS,
tunnel_connect_timeout_ms: 2_000, tunnel_connect_timeout_ms: 2_000,
tunnel_ipv4_only: false, tunnel_ipv4_only: false,
tunnel_ipv6_only: false, tunnel_ipv6_only: false,
+233 -17
View File
@@ -25,7 +25,7 @@ use crate::upstream_client;
use super::protocol::{ use super::protocol::{
compress_payload, decompress_if_gzip, flags, raw_payload, Frame as TunnelFrame, MsgType, compress_payload, decompress_if_gzip, flags, raw_payload, Frame as TunnelFrame, MsgType,
RequestMeta, ResponseMeta, RequestMeta, ResetStreamPayload, ResponseMeta,
}; };
use super::writer::FrameSender; use super::writer::FrameSender;
@@ -35,10 +35,101 @@ const MAX_CHUNK_SIZE: usize = 32 * 1024;
/// Timeout for sending a single frame to the writer channel. /// Timeout for sending a single frame to the writer channel.
/// Control frames are allowed a short wait; body frames fail fast. /// Control frames are allowed a short wait; body frames fail fast.
const CONTROL_FRAME_SEND_TIMEOUT: Duration = Duration::from_millis(250); const CONTROL_FRAME_SEND_TIMEOUT: Duration = Duration::from_millis(250);
const FLOW_CONTROL_WAIT_TIMEOUT: Duration = Duration::from_secs(5);
const SLOW_STREAM_LOG_THRESHOLD: Duration = Duration::from_secs(2); const SLOW_STREAM_LOG_THRESHOLD: Duration = Duration::from_secs(2);
const SUCCESS_LOG_SAMPLE_MODULO: u32 = 256; const SUCCESS_LOG_SAMPLE_MODULO: u32 = 256;
const REQUEST_BODY_SPOOL_QUEUE_CAPACITY: usize = 64; const REQUEST_BODY_SPOOL_QUEUE_CAPACITY: usize = 64;
#[derive(Debug)]
pub(crate) struct StreamSendWindow {
available: Mutex<u64>,
notify: Notify,
}
impl StreamSendWindow {
pub(crate) fn new(initial_window_bytes: u32) -> Self {
Self {
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
notify: Notify::new(),
}
}
pub(crate) fn add_credit(&self, delta_bytes: u32) {
if delta_bytes == 0 {
return;
}
let mut available = self.available.lock().expect("stream window lock poisoned");
*available = available.saturating_add(u64::from(delta_bytes));
drop(available);
self.notify.notify_waiters();
}
async fn acquire(&self, bytes: usize, timeout: Duration) -> Result<Duration, ()> {
if bytes == 0 {
return Ok(Duration::ZERO);
}
let requested = bytes as u64;
let started_at = Instant::now();
loop {
{
let mut available = self.available.lock().expect("stream window lock poisoned");
if *available >= requested {
*available -= requested;
return Ok(started_at.elapsed());
}
}
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(());
};
if tokio::time::timeout(remaining, self.notify.notified())
.await
.is_err()
{
return Err(());
}
}
}
}
fn stream_reset_message(frame: &TunnelFrame) -> String {
if frame.msg_type == MsgType::ResetStream {
if let Ok(payload) = serde_json::from_slice::<ResetStreamPayload>(&frame.payload) {
return payload.reason;
}
}
String::from_utf8(frame.payload.to_vec())
.unwrap_or_else(|_| "client cancelled request body".to_string())
}
fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) {
if bytes == 0 {
return;
}
let delta = bytes.min(u32::MAX as usize) as u32;
if frame_tx
.try_send(TunnelFrame::new(
stream_id,
MsgType::WindowUpdate,
0,
Bytes::from(
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
delta_bytes: delta,
})
.expect("window update payload should serialize"),
),
))
.is_err()
{
warn!(
stream_id,
delta_bytes = delta,
"writer channel full, WINDOW_UPDATE dropped"
);
}
}
/// Minimum allowed upstream request timeout (milliseconds). /// Minimum allowed upstream request timeout (milliseconds).
const MIN_TIMEOUT_MS: u64 = 1; const MIN_TIMEOUT_MS: u64 = 1;
/// Maximum allowed upstream request timeout (milliseconds). /// Maximum allowed upstream request timeout (milliseconds).
@@ -491,10 +582,12 @@ fn buffered_request_body(body: Bytes) -> upstream_client::UpstreamRequestBody {
// longer coupled to upstream body polling. Redirect replay still reuses a full // longer coupled to upstream body polling. Redirect replay still reuses a full
// in-memory copy when the request body completes within budget. // in-memory copy when the request body completes within budget.
fn prepare_request_body( fn prepare_request_body(
stream_id: u32,
body_rx: mpsc::Receiver<TunnelFrame>, body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>, body_size: Arc<AtomicUsize>,
deadline: Instant, deadline: Instant,
replay_budget_bytes: usize, replay_budget_bytes: usize,
frame_tx: FrameSender,
) -> PreparedRequestBody { ) -> PreparedRequestBody {
let (spool_tx, spool_rx) = mpsc::channel(REQUEST_BODY_SPOOL_QUEUE_CAPACITY); let (spool_tx, spool_rx) = mpsc::channel(REQUEST_BODY_SPOOL_QUEUE_CAPACITY);
let replay_state = if replay_budget_bytes == 0 { let replay_state = if replay_budget_bytes == 0 {
@@ -508,11 +601,13 @@ fn prepare_request_body(
}; };
tokio::spawn(spool_request_body( tokio::spawn(spool_request_body(
stream_id,
body_rx, body_rx,
spool_tx, spool_tx,
replay_state, replay_state,
body_size, body_size,
deadline, deadline,
frame_tx,
)); ));
PreparedRequestBody { PreparedRequestBody {
@@ -537,10 +632,12 @@ fn prepare_bodyless_request_body(
} }
async fn collect_request_body_for_replay( async fn collect_request_body_for_replay(
stream_id: u32,
mut body_rx: mpsc::Receiver<TunnelFrame>, mut body_rx: mpsc::Receiver<TunnelFrame>,
body_size: Arc<AtomicUsize>, body_size: Arc<AtomicUsize>,
deadline: Instant, deadline: Instant,
replay_budget_bytes: usize, replay_budget_bytes: usize,
frame_tx: &FrameSender,
) -> Result<Bytes, String> { ) -> Result<Bytes, String> {
let mut body = BytesMut::new(); let mut body = BytesMut::new();
@@ -565,6 +662,7 @@ async fn collect_request_body_for_replay(
)); ));
} }
body_size.fetch_add(payload.len(), Ordering::Relaxed); body_size.fetch_add(payload.len(), Ordering::Relaxed);
try_send_window_update(frame_tx, stream_id, payload.len());
body.extend_from_slice(&payload); body.extend_from_slice(&payload);
} }
@@ -572,9 +670,8 @@ async fn collect_request_body_for_replay(
return Ok(body.freeze()); return Ok(body.freeze());
} }
} }
MsgType::StreamError => { MsgType::StreamError | MsgType::ResetStream => {
return Err(String::from_utf8(frame.payload.to_vec()) return Err(stream_reset_message(&frame));
.unwrap_or_else(|_| "client cancelled request body".to_string()));
} }
MsgType::StreamEnd => return Ok(body.freeze()), MsgType::StreamEnd => return Ok(body.freeze()),
_ => continue, _ => continue,
@@ -648,11 +745,13 @@ fn timeout_duration_from_legacy_secs(secs: u64) -> Duration {
} }
async fn spool_request_body( async fn spool_request_body(
stream_id: u32,
mut body_rx: mpsc::Receiver<TunnelFrame>, mut body_rx: mpsc::Receiver<TunnelFrame>,
mut spool_tx: mpsc::Sender<SpoolBodyEvent>, mut spool_tx: mpsc::Sender<SpoolBodyEvent>,
replay_state: Option<Arc<RequestBodyReplayState>>, replay_state: Option<Arc<RequestBodyReplayState>>,
body_size: Arc<AtomicUsize>, body_size: Arc<AtomicUsize>,
deadline: Instant, deadline: Instant,
frame_tx: FrameSender,
) { ) {
loop { loop {
let frame = match recv_body_frame_with_deadline(&mut body_rx, deadline).await { let frame = match recv_body_frame_with_deadline(&mut body_rx, deadline).await {
@@ -692,6 +791,7 @@ async fn spool_request_body(
if !payload.is_empty() { if !payload.is_empty() {
body_size.fetch_add(payload.len(), Ordering::Relaxed); body_size.fetch_add(payload.len(), Ordering::Relaxed);
try_send_window_update(&frame_tx, stream_id, payload.len());
if let Some(state) = &replay_state { if let Some(state) = &replay_state {
state.push_chunk(payload.clone()); state.push_chunk(payload.clone());
} }
@@ -714,9 +814,8 @@ async fn spool_request_body(
return; return;
} }
} }
MsgType::StreamError => { MsgType::StreamError | MsgType::ResetStream => {
let message = String::from_utf8(frame.payload.to_vec()) let message = stream_reset_message(&frame);
.unwrap_or_else(|_| "client cancelled request body".to_string());
if let Some(state) = &replay_state { if let Some(state) = &replay_state {
state.fail(message.clone()); state.fail(message.clone());
} }
@@ -932,6 +1031,40 @@ async fn execute_upstream_request(
}) })
} }
async fn acquire_response_credit(
response_window: &StreamSendWindow,
frame_tx: &FrameSender,
stream_id: u32,
bytes: usize,
) -> bool {
match response_window
.acquire(bytes, FLOW_CONTROL_WAIT_TIMEOUT)
.await
{
Ok(waited) => {
if waited > Duration::from_millis(1) {
debug!(
stream_id,
bytes,
waited_ms = waited.as_millis() as u64,
"waited for tunnel response flow-control credit"
);
}
true
}
Err(()) => {
warn!(
stream_id,
bytes,
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
"response flow-control window timeout"
);
send_reset_stream(frame_tx, stream_id, "response_flow_control_timeout").await;
false
}
}
}
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
async fn relay_upstream_response<B>( async fn relay_upstream_response<B>(
server: &ServerContext, server: &ServerContext,
@@ -939,6 +1072,7 @@ async fn relay_upstream_response<B>(
method: &hyper::Method, method: &hyper::Method,
request_url: &url::Url, request_url: &url::Url,
frame_tx: &FrameSender, frame_tx: &FrameSender,
response_window: &StreamSendWindow,
response: hyper::Response<B>, response: hyper::Response<B>,
total_dns_ms: u64, total_dns_ms: u64,
total_elapsed: Duration, total_elapsed: Duration,
@@ -1068,6 +1202,11 @@ where
Ok(chunk) => { Ok(chunk) => {
if chunk.len() <= MAX_CHUNK_SIZE { if chunk.len() <= MAX_CHUNK_SIZE {
let (payload, extra_flags) = raw_payload(chunk); let (payload, extra_flags) = raw_payload(chunk);
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
.await
{
return Some(total_elapsed);
}
if !send_frame( if !send_frame(
frame_tx, frame_tx,
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload), TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
@@ -1094,6 +1233,16 @@ where
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len()); let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
let slice = chunk.slice(offset..end); let slice = chunk.slice(offset..end);
let (payload, extra_flags) = raw_payload(slice); let (payload, extra_flags) = raw_payload(slice);
if !acquire_response_credit(
response_window,
frame_tx,
stream_id,
payload.len(),
)
.await
{
return Some(total_elapsed);
}
if !send_frame( if !send_frame(
frame_tx, frame_tx,
TunnelFrame::new( TunnelFrame::new(
@@ -1213,6 +1362,7 @@ pub async fn handle_stream(
meta: RequestMeta, meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>, body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: FrameSender, frame_tx: FrameSender,
response_window: Arc<StreamSendWindow>,
) { ) {
let request_method = parse_request_method(&meta.method); let request_method = parse_request_method(&meta.method);
let request_url = url::Url::parse(&meta.url).ok(); let request_url = url::Url::parse(&meta.url).ok();
@@ -1244,8 +1394,17 @@ pub async fn handle_stream(
server.active_connections.fetch_add(1, Ordering::Release); server.active_connections.fetch_add(1, Ordering::Release);
let connect_elapsed = let connect_elapsed = handle_stream_inner(
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx, permit).await; &state,
&server,
stream_id,
meta,
body_rx,
&frame_tx,
response_window.as_ref(),
permit,
)
.await;
server.active_connections.fetch_sub(1, Ordering::Release); server.active_connections.fetch_sub(1, Ordering::Release);
if let Some(d) = connect_elapsed { if let Some(d) = connect_elapsed {
@@ -1264,18 +1423,21 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
); );
if is_body_frame { if is_body_frame {
match tx.try_send(frame) { match tokio::time::timeout(FLOW_CONTROL_WAIT_TIMEOUT, tx.send(frame)).await {
Ok(()) => true, Ok(Ok(())) => true,
Err(QueueSendError::Full(_)) => { Ok(Err(QueueSendError::Closed(_))) | Err(_) => {
warn!( warn!(
stream_id, stream_id,
msg_type = ?msg_type, msg_type = ?msg_type,
flags = flags, flags = flags,
"writer channel full for body frame, abandoning stream" timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
"writer channel stalled for body frame, abandoning stream"
); );
false false
} }
Err(QueueSendError::Closed(_)) => false, Ok(Err(QueueSendError::Full(_))) => {
unreachable!("bounded queue send should not report full")
}
} }
} else { } else {
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await { match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
@@ -1304,6 +1466,7 @@ async fn handle_stream_inner(
meta: RequestMeta, meta: RequestMeta,
body_rx: mpsc::Receiver<TunnelFrame>, body_rx: mpsc::Receiver<TunnelFrame>,
frame_tx: &FrameSender, frame_tx: &FrameSender,
response_window: &StreamSendWindow,
mut admission_permit: Option<AdmissionPermit>, mut admission_permit: Option<AdmissionPermit>,
) -> Option<Duration> { ) -> Option<Duration> {
let mut current_method: hyper::Method = parse_request_method(&meta.method); let mut current_method: hyper::Method = parse_request_method(&meta.method);
@@ -1357,10 +1520,12 @@ async fn handle_stream_inner(
}; };
let mut prepared_body = if can_buffer_redirect_body { let mut prepared_body = if can_buffer_redirect_body {
let buffered_body = match collect_request_body_for_replay( let buffered_body = match collect_request_body_for_replay(
stream_id,
body_rx, body_rx,
Arc::clone(&request_body_size), Arc::clone(&request_body_size),
first_byte_deadline, first_byte_deadline,
replay_budget_bytes, replay_budget_bytes,
frame_tx,
) )
.await .await
{ {
@@ -1388,10 +1553,12 @@ async fn handle_stream_inner(
} }
} else if request_has_body { } else if request_has_body {
prepare_request_body( prepare_request_body(
stream_id,
body_rx, body_rx,
Arc::clone(&request_body_size), Arc::clone(&request_body_size),
first_byte_deadline, first_byte_deadline,
0, 0,
frame_tx.clone(),
) )
} else { } else {
prepare_bodyless_request_body(body_rx, follow_redirects) prepare_bodyless_request_body(body_rx, follow_redirects)
@@ -1472,6 +1639,7 @@ async fn handle_stream_inner(
&current_method, &current_method,
&current_url, &current_url,
frame_tx, frame_tx,
response_window,
response_ctx.response, response_ctx.response,
total_dns_ms, total_dns_ms,
overall_start.elapsed(), overall_start.elapsed(),
@@ -1512,6 +1680,7 @@ async fn handle_stream_inner(
&current_method, &current_method,
&current_url, &current_url,
frame_tx, frame_tx,
response_window,
response_ctx.response, response_ctx.response,
total_dns_ms, total_dns_ms,
overall_start.elapsed(), overall_start.elapsed(),
@@ -1568,6 +1737,7 @@ async fn handle_stream_inner(
&current_method, &current_method,
&current_url, &current_url,
frame_tx, frame_tx,
response_window,
response_ctx.response, response_ctx.response,
total_dns_ms, total_dns_ms,
overall_start.elapsed(), overall_start.elapsed(),
@@ -1596,6 +1766,18 @@ async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
.await; .await;
} }
async fn send_reset_stream(tx: &FrameSender, stream_id: u32, reason: &str) {
let payload = serde_json::to_vec(&ResetStreamPayload {
reason: reason.to_string(),
})
.expect("reset stream payload should serialize");
let _ = send_frame(
tx,
TunnelFrame::new(stream_id, MsgType::ResetStream, 0, Bytes::from(payload)),
)
.await;
}
#[cfg(test)] #[cfg(test)]
fn build_streaming_request_body( fn build_streaming_request_body(
body_rx: mpsc::Receiver<TunnelFrame>, body_rx: mpsc::Receiver<TunnelFrame>,
@@ -1676,9 +1858,8 @@ fn build_prefixed_request_body(
(body_rx, body_size, end_stream), (body_rx, body_size, end_stream),
)); ));
} }
MsgType::StreamError => { MsgType::StreamError | MsgType::ResetStream => {
let message = String::from_utf8(frame.payload.to_vec()) let message = stream_reset_message(&frame);
.unwrap_or_else(|_| "client cancelled request body".to_string());
return Some((Err(io::Error::other(message)), (body_rx, body_size, true))); return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
} }
MsgType::StreamEnd => return None, MsgType::StreamEnd => return None,
@@ -1820,12 +2001,15 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn prepare_request_body_streams_immediately_and_replays_after_completion() { async fn prepare_request_body_streams_immediately_and_replays_after_completion() {
let (tx, rx) = mpsc::channel(4); let (tx, rx) = mpsc::channel(4);
let (frame_tx, sent, writer_handle) = spawn_test_writer();
let body_size = Arc::new(AtomicUsize::new(0)); let body_size = Arc::new(AtomicUsize::new(0));
let prepared = prepare_request_body( let prepared = prepare_request_body(
1,
rx, rx,
Arc::clone(&body_size), Arc::clone(&body_size),
Instant::now() + Duration::from_secs(1), Instant::now() + Duration::from_secs(1),
1024, 1024,
frame_tx.clone(),
); );
let mut body = prepared let mut body = prepared
.first_request_body .first_request_body
@@ -1888,6 +2072,19 @@ mod tests {
); );
assert!(replay.frame().await.is_none()); assert!(replay.frame().await.is_none());
assert_eq!(body_size.load(Ordering::Relaxed), 11); assert_eq!(body_size.load(Ordering::Relaxed), 11);
let window_update_bytes = collect_emitted_frames(frame_tx, sent, writer_handle)
.await
.into_iter()
.filter(|frame| frame.msg_type == MsgType::WindowUpdate)
.filter_map(|frame| {
serde_json::from_slice::<aether_contracts::tunnel::WindowUpdatePayload>(
&frame.payload,
)
.ok()
})
.map(|payload| payload.delta_bytes as usize)
.sum::<usize>();
assert_eq!(window_update_bytes, 11);
} }
#[test] #[test]
@@ -2084,6 +2281,7 @@ mod tests {
meta, meta,
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await; let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2145,6 +2343,7 @@ mod tests {
meta, meta,
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await; let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2179,6 +2378,7 @@ mod tests {
.status(StatusCode::OK) .status(StatusCode::OK)
.body(body) .body(body)
.expect("response"); .expect("response");
let response_window = test_response_window();
relay_upstream_response( relay_upstream_response(
&server, &server,
@@ -2186,6 +2386,7 @@ mod tests {
&hyper::Method::GET, &hyper::Method::GET,
&request_url, &request_url,
&frame_tx, &frame_tx,
response_window.as_ref(),
response, response,
0, 0,
Duration::ZERO, Duration::ZERO,
@@ -2222,6 +2423,7 @@ mod tests {
.status(StatusCode::OK) .status(StatusCode::OK)
.body(body) .body(body)
.expect("response"); .expect("response");
let response_window = test_response_window();
relay_upstream_response( relay_upstream_response(
&server, &server,
@@ -2229,6 +2431,7 @@ mod tests {
&hyper::Method::GET, &hyper::Method::GET,
&request_url, &request_url,
&frame_tx, &frame_tx,
response_window.as_ref(),
response, response,
0, 0,
Duration::ZERO, Duration::ZERO,
@@ -2329,6 +2532,7 @@ mod tests {
meta, meta,
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await; let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2393,6 +2597,7 @@ mod tests {
meta, meta,
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await; let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2477,6 +2682,7 @@ mod tests {
meta, meta,
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
let result = collect_stream_result(frame_tx, sent, writer_handle).await; let result = collect_stream_result(frame_tx, sent, writer_handle).await;
@@ -2515,6 +2721,7 @@ mod tests {
sample_request_meta(), sample_request_meta(),
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
@@ -2561,6 +2768,7 @@ mod tests {
sample_request_meta(), sample_request_meta(),
body_rx, body_rx,
frame_tx.clone(), frame_tx.clone(),
test_response_window(),
) )
.await; .await;
@@ -2729,6 +2937,10 @@ mod tests {
tunnel_reconnect_max_ms: 30_000, tunnel_reconnect_max_ms: 30_000,
tunnel_ping_interval_ms: 15_000, tunnel_ping_interval_ms: 15_000,
tunnel_max_streams: Some(8), tunnel_max_streams: Some(8),
tunnel_profile: crate::config::TunnelProfileArg::Lite,
tunnel_stream_initial_window_bytes:
crate::config::DEFAULT_TUNNEL_STREAM_INITIAL_WINDOW_BYTES,
tunnel_drain_deadline_ms: crate::config::DEFAULT_TUNNEL_DRAIN_DEADLINE_MS,
tunnel_connect_timeout_ms: 15_000, tunnel_connect_timeout_ms: 15_000,
tunnel_ipv4_only: false, tunnel_ipv4_only: false,
tunnel_ipv6_only: false, tunnel_ipv6_only: false,
@@ -2793,6 +3005,10 @@ mod tests {
(frame_tx, sent, handle) (frame_tx, sent, handle)
} }
fn test_response_window() -> Arc<StreamSendWindow> {
Arc::new(StreamSendWindow::new(u32::MAX))
}
struct StreamResult { struct StreamResult {
response: Option<ResponseMeta>, response: Option<ResponseMeta>,
body: Bytes, body: Bytes,
+7 -1
View File
@@ -188,7 +188,13 @@ fn classify_frame_priority(frame: &Frame) -> FramePriority {
| MsgType::Pong | MsgType::Pong
| MsgType::GoAway | MsgType::GoAway
| MsgType::HeartbeatData | MsgType::HeartbeatData
| MsgType::HeartbeatAck => FramePriority::High, | MsgType::HeartbeatAck
| MsgType::Hello
| MsgType::Settings
| MsgType::WindowUpdate
| MsgType::ResetStream
| MsgType::ConnectionClose
| MsgType::LoadReport => FramePriority::High,
MsgType::RequestHeaders MsgType::RequestHeaders
| MsgType::RequestBody | MsgType::RequestBody
| MsgType::ResponseBody | MsgType::ResponseBody
+157 -8
View File
@@ -10,8 +10,8 @@ pub const TUNNEL_RELAY_FORWARDED_BY_HEADER: &str = "x-aether-tunnel-forwarded-by
pub const TUNNEL_RELAY_OWNER_INSTANCE_HEADER: &str = "x-aether-tunnel-owner-instance-id"; pub const TUNNEL_RELAY_OWNER_INSTANCE_HEADER: &str = "x-aether-tunnel-owner-instance-id";
pub const TUNNEL_PROTOCOL_VERSION_HEADER: &str = "x-aether-tunnel-protocol-version"; pub const TUNNEL_PROTOCOL_VERSION_HEADER: &str = "x-aether-tunnel-protocol-version";
pub const TUNNEL_NODE_NAME_B64_HEADER: &str = "x-aether-tunnel-node-name-b64"; pub const TUNNEL_NODE_NAME_B64_HEADER: &str = "x-aether-tunnel-node-name-b64";
pub const CURRENT_TUNNEL_PROTOCOL_VERSION: u8 = 2; pub const CURRENT_TUNNEL_PROTOCOL_VERSION: u8 = 3;
pub const CURRENT_TUNNEL_PROTOCOL_VERSION_STR: &str = "2"; pub const CURRENT_TUNNEL_PROTOCOL_VERSION_STR: &str = "3";
pub mod flags { pub mod flags {
pub const END_STREAM: u8 = 0x01; pub const END_STREAM: u8 = 0x01;
@@ -33,6 +33,12 @@ pub enum MsgType {
GoAway = 0x12, GoAway = 0x12,
HeartbeatData = 0x13, HeartbeatData = 0x13,
HeartbeatAck = 0x14, HeartbeatAck = 0x14,
Hello = 0x15,
Settings = 0x16,
WindowUpdate = 0x17,
ResetStream = 0x18,
ConnectionClose = 0x19,
LoadReport = 0x1a,
} }
impl MsgType { impl MsgType {
@@ -49,6 +55,12 @@ impl MsgType {
GOAWAY => Some(Self::GoAway), GOAWAY => Some(Self::GoAway),
HEARTBEAT_DATA => Some(Self::HeartbeatData), HEARTBEAT_DATA => Some(Self::HeartbeatData),
HEARTBEAT_ACK => Some(Self::HeartbeatAck), HEARTBEAT_ACK => Some(Self::HeartbeatAck),
HELLO => Some(Self::Hello),
SETTINGS => Some(Self::Settings),
WINDOW_UPDATE => Some(Self::WindowUpdate),
RESET_STREAM => Some(Self::ResetStream),
CONNECTION_CLOSE => Some(Self::ConnectionClose),
LOAD_REPORT => Some(Self::LoadReport),
_ => None, _ => None,
} }
} }
@@ -65,6 +77,12 @@ pub const PONG: u8 = MsgType::Pong as u8;
pub const GOAWAY: u8 = MsgType::GoAway as u8; pub const GOAWAY: u8 = MsgType::GoAway as u8;
pub const HEARTBEAT_DATA: u8 = MsgType::HeartbeatData as u8; pub const HEARTBEAT_DATA: u8 = MsgType::HeartbeatData as u8;
pub const HEARTBEAT_ACK: u8 = MsgType::HeartbeatAck as u8; pub const HEARTBEAT_ACK: u8 = MsgType::HeartbeatAck as u8;
pub const HELLO: u8 = MsgType::Hello as u8;
pub const SETTINGS: u8 = MsgType::Settings as u8;
pub const WINDOW_UPDATE: u8 = MsgType::WindowUpdate as u8;
pub const RESET_STREAM: u8 = MsgType::ResetStream as u8;
pub const CONNECTION_CLOSE: u8 = MsgType::ConnectionClose as u8;
pub const LOAD_REPORT: u8 = MsgType::LoadReport as u8;
pub const FLAG_END_STREAM: u8 = flags::END_STREAM; pub const FLAG_END_STREAM: u8 = flags::END_STREAM;
pub const FLAG_GZIP_COMPRESSED: u8 = flags::GZIP_COMPRESSED; pub const FLAG_GZIP_COMPRESSED: u8 = flags::GZIP_COMPRESSED;
pub const FLAG_ENCRYPTED: u8 = flags::ENCRYPTED; pub const FLAG_ENCRYPTED: u8 = flags::ENCRYPTED;
@@ -246,6 +264,54 @@ pub struct ResponseMeta {
pub headers: Vec<(String, String)>, pub headers: Vec<(String, String)>,
} }
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct HelloPayload {
pub protocol_version: u8,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub capabilities: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub session_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub replica_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct SettingsPayload {
pub initial_stream_window_bytes: u32,
pub min_window_update_bytes: u32,
pub drain_deadline_ms: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct WindowUpdatePayload {
pub delta_bytes: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ResetStreamPayload {
pub reason: String,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct GoAwayPayload {
pub last_accepted_stream_id: u32,
pub drain_deadline_ms: u64,
pub reason: String,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ConnectionClosePayload {
pub reason: String,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct LoadReportPayload {
pub active_streams: u32,
pub queue_depth: u32,
pub queue_capacity: u32,
pub health_score: u8,
}
pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> { pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len()); let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
buf.extend_from_slice(&stream_id.to_be_bytes()); buf.extend_from_slice(&stream_id.to_be_bytes());
@@ -260,6 +326,14 @@ pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes()) encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes())
} }
pub fn encode_reset_stream(stream_id: u32, reason: &str) -> Vec<u8> {
let payload = serde_json::to_vec(&ResetStreamPayload {
reason: reason.to_string(),
})
.expect("reset stream payload should serialize");
encode_frame(stream_id, RESET_STREAM, 0, &payload)
}
pub fn encode_ping() -> Vec<u8> { pub fn encode_ping() -> Vec<u8> {
encode_frame(0, PING, 0, &[]) encode_frame(0, PING, 0, &[])
} }
@@ -272,6 +346,51 @@ pub fn encode_goaway() -> Vec<u8> {
encode_frame(0, GOAWAY, 0, &[]) encode_frame(0, GOAWAY, 0, &[])
} }
pub fn encode_goaway_v3(
last_accepted_stream_id: u32,
drain_deadline_ms: u64,
reason: &str,
) -> Vec<u8> {
let payload = serde_json::to_vec(&GoAwayPayload {
last_accepted_stream_id,
drain_deadline_ms,
reason: reason.to_string(),
})
.expect("goaway payload should serialize");
encode_frame(0, GOAWAY, 0, &payload)
}
pub fn encode_hello(payload: &HelloPayload) -> Vec<u8> {
encode_json_control(HELLO, payload)
}
pub fn encode_settings(payload: &SettingsPayload) -> Vec<u8> {
encode_json_control(SETTINGS, payload)
}
pub fn encode_window_update(stream_id: u32, delta_bytes: u32) -> Vec<u8> {
let payload = serde_json::to_vec(&WindowUpdatePayload { delta_bytes })
.expect("window update payload should serialize");
encode_frame(stream_id, WINDOW_UPDATE, 0, &payload)
}
pub fn encode_connection_close(reason: &str) -> Vec<u8> {
let payload = serde_json::to_vec(&ConnectionClosePayload {
reason: reason.to_string(),
})
.expect("connection close payload should serialize");
encode_frame(0, CONNECTION_CLOSE, 0, &payload)
}
pub fn encode_load_report(payload: &LoadReportPayload) -> Vec<u8> {
encode_json_control(LOAD_REPORT, payload)
}
fn encode_json_control<T: serde::Serialize>(msg_type: u8, payload: &T) -> Vec<u8> {
let payload = serde_json::to_vec(payload).expect("tunnel control payload should serialize");
encode_frame(0, msg_type, 0, &payload)
}
#[inline] #[inline]
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> { pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
let payload_len = header.payload_len as usize; let payload_len = header.payload_len as usize;
@@ -341,10 +460,11 @@ fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
compress_payload, decode_payload, encode_frame, encode_ping, raw_payload, Frame, compress_payload, decode_payload, encode_frame, encode_goaway_v3, encode_ping,
FrameHeader, MsgType, RequestMeta, CURRENT_TUNNEL_PROTOCOL_VERSION, encode_reset_stream, encode_window_update, raw_payload, Frame, FrameHeader, GoAwayPayload,
CURRENT_TUNNEL_PROTOCOL_VERSION_STR, FLAG_GZIP_COMPRESSED, REQUEST_HEADERS, MsgType, RequestMeta, ResetStreamPayload, WindowUpdatePayload,
TUNNEL_PROTOCOL_VERSION_HEADER, CURRENT_TUNNEL_PROTOCOL_VERSION, CURRENT_TUNNEL_PROTOCOL_VERSION_STR, FLAG_GZIP_COMPRESSED,
REQUEST_HEADERS, TUNNEL_PROTOCOL_VERSION_HEADER,
}; };
use bytes::Bytes; use bytes::Bytes;
@@ -407,7 +527,36 @@ mod tests {
TUNNEL_PROTOCOL_VERSION_HEADER, TUNNEL_PROTOCOL_VERSION_HEADER,
"x-aether-tunnel-protocol-version" "x-aether-tunnel-protocol-version"
); );
assert_eq!(CURRENT_TUNNEL_PROTOCOL_VERSION, 2); assert_eq!(CURRENT_TUNNEL_PROTOCOL_VERSION, 3);
assert_eq!(CURRENT_TUNNEL_PROTOCOL_VERSION_STR, "2"); assert_eq!(CURRENT_TUNNEL_PROTOCOL_VERSION_STR, "3");
}
#[test]
fn v3_control_frames_round_trip_json_payloads() {
let reset = encode_reset_stream(9, "request body window exhausted");
let reset_header = FrameHeader::parse(&reset).expect("reset header");
assert_eq!(reset_header.msg_type, super::RESET_STREAM);
let reset_payload = decode_payload(&reset, &reset_header).expect("reset payload");
let reset_payload: ResetStreamPayload =
serde_json::from_slice(&reset_payload).expect("reset json");
assert_eq!(reset_payload.reason, "request body window exhausted");
let window = encode_window_update(9, 1024 * 1024);
let window_header = FrameHeader::parse(&window).expect("window header");
assert_eq!(window_header.msg_type, super::WINDOW_UPDATE);
let window_payload = decode_payload(&window, &window_header).expect("window payload");
let window_payload: WindowUpdatePayload =
serde_json::from_slice(&window_payload).expect("window json");
assert_eq!(window_payload.delta_bytes, 1024 * 1024);
let goaway = encode_goaway_v3(42, 30_000, "rolling restart");
let goaway_header = FrameHeader::parse(&goaway).expect("goaway header");
assert_eq!(goaway_header.msg_type, super::GOAWAY);
let goaway_payload = decode_payload(&goaway, &goaway_header).expect("goaway payload");
let goaway_payload: GoAwayPayload =
serde_json::from_slice(&goaway_payload).expect("goaway json");
assert_eq!(goaway_payload.last_accepted_stream_id, 42);
assert_eq!(goaway_payload.drain_deadline_ms, 30_000);
assert_eq!(goaway_payload.reason, "rolling restart");
} }
} }
@@ -660,6 +660,29 @@ async fn connect_protocol_peer(
let (socket, _response) = tokio_tungstenite::connect_async(request).await?; let (socket, _response) = tokio_tungstenite::connect_async(request).await?;
let (mut sink, mut stream) = socket.split(); let (mut sink, mut stream) = socket.split();
sink.send(Message::Binary(
protocol::encode_hello(&protocol::HelloPayload {
protocol_version: aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION,
capabilities: vec![
"flow-control".to_string(),
"reset-stream".to_string(),
"graceful-drain".to_string(),
],
session_id: Some("capacity-curve-session".to_string()),
replica_id: Some("capacity-curve-replica".to_string()),
})
.into(),
))
.await?;
sink.send(Message::Binary(
protocol::encode_settings(&protocol::SettingsPayload {
initial_stream_window_bytes: 4 * 1024 * 1024,
min_window_update_bytes: 1024 * 1024,
drain_deadline_ms: 30_000,
})
.into(),
))
.await?;
Ok(tokio::spawn(async move { Ok(tokio::spawn(async move {
while let Some(message) = stream.next().await { while let Some(message) = stream.next().await {
let Ok(message) = message else { let Ok(message) = message else {
@@ -707,7 +730,17 @@ where
let payload = protocol::decode_payload(&data, &header).unwrap_or_default(); let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload); let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload);
} }
protocol::REQUEST_BODY if header.flags & protocol::FLAG_END_STREAM != 0 => { protocol::REQUEST_BODY => {
let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
if !payload.is_empty() {
sink.send(Message::Binary(
protocol::encode_window_update(header.stream_id, payload.len() as u32).into(),
))
.await?;
}
if header.flags & protocol::FLAG_END_STREAM == 0 {
return Ok(());
}
tokio::time::sleep(hold).await; tokio::time::sleep(hold).await;
let response_meta = protocol::ResponseMeta { let response_meta = protocol::ResponseMeta {
status: 200, status: 200,
@@ -1,11 +1,12 @@
use std::collections::BTreeMap;
use std::path::PathBuf; use std::path::PathBuf;
use std::time::Duration; use std::time::Duration;
use aether_gateway::tunnel_protocol as protocol; use aether_gateway::tunnel_protocol as protocol;
use aether_testkit::{ use aether_testkit::{
fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, run_http_load_probe, fetch_prometheus_samples, find_metric_value_u64, init_test_runtime_for, run_http_load_probe,
HttpLoadProbeConfig, HttpLoadProbeResponseMode, HttpLoadProbeResult, TunnelHarness, HttpLoadProbeConfig, HttpLoadProbeResponseMode, HttpLoadProbeResult, PrometheusSample,
TunnelHarnessConfig, TunnelHarness, TunnelHarnessConfig,
}; };
use futures_util::{SinkExt, StreamExt}; use futures_util::{SinkExt, StreamExt};
use reqwest::Method; use reqwest::Method;
@@ -20,6 +21,11 @@ const TUNNEL_RELAY_PATH_PREFIX: &str = "/api/internal/tunnel/relay";
struct GatewayTunnelBaselineConfig { struct GatewayTunnelBaselineConfig {
total_requests: usize, total_requests: usize,
concurrency: usize, concurrency: usize,
request_body_bytes: usize,
tunnel_connections: usize,
outbound_queue_capacity: usize,
close_one_connection_after: Option<Duration>,
require_acceptance: bool,
timeout: Duration, timeout: Duration,
output_path: Option<PathBuf>, output_path: Option<PathBuf>,
} }
@@ -29,6 +35,11 @@ impl Default for GatewayTunnelBaselineConfig {
Self { Self {
total_requests: 200, total_requests: 200,
concurrency: 20, concurrency: 20,
request_body_bytes: 6 * 1024 * 1024,
tunnel_connections: 4,
outbound_queue_capacity: 512,
close_one_connection_after: None,
require_acceptance: false,
timeout: Duration::from_secs(10), timeout: Duration::from_secs(10),
output_path: None, output_path: None,
} }
@@ -38,8 +49,21 @@ impl Default for GatewayTunnelBaselineConfig {
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
struct GatewayTunnelBaselineReport { struct GatewayTunnelBaselineReport {
suite: &'static str, suite: &'static str,
config: GatewayTunnelEffectiveConfig,
scenario: HttpLoadProbeResult, scenario: HttpLoadProbeResult,
tunnel_metrics: TunnelMetricsSnapshot, tunnel_metrics: TunnelMetricsSnapshot,
acceptance: AcceptanceReport,
}
#[derive(Debug, Serialize)]
struct GatewayTunnelEffectiveConfig {
total_requests: usize,
concurrency: usize,
request_body_bytes: usize,
tunnel_connections: usize,
outbound_queue_capacity: usize,
close_one_connection_after_ms: Option<u64>,
timeout_ms: u64,
} }
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
@@ -54,8 +78,28 @@ struct TunnelMetricsSnapshot {
proxy_connection_congested_total: u64, proxy_connection_congested_total: u64,
proxy_connection_write_latency_last_us_max: u64, proxy_connection_write_latency_last_us_max: u64,
proxy_connection_write_latency_ewma_us_max: u64, proxy_connection_write_latency_ewma_us_max: u64,
body_backpressure_total: u64,
flow_window_blocked_ms: u64,
connection_health_score: u64,
stream_reset_total: u64,
stream_reset_reasons: BTreeMap<String, u64>,
drain_total: u64,
drain_reasons: BTreeMap<String, u64>,
scheduler_selected_conn_total: u64,
proxy_connections_protocol_v1: u64, proxy_connections_protocol_v1: u64,
proxy_connections_protocol_v2: u64, proxy_connections_protocol_v2: u64,
proxy_connections_protocol_v3: u64,
}
#[derive(Debug, Serialize)]
struct AcceptanceReport {
required: bool,
passed: bool,
success_rate_bps: u64,
min_success_rate_bps: u64,
congestion_free: bool,
no_queue_full_rejections: bool,
reasons: Vec<String>,
} }
#[tokio::main] #[tokio::main]
@@ -77,8 +121,21 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run_suite( async fn run_suite(
config: &GatewayTunnelBaselineConfig, config: &GatewayTunnelBaselineConfig,
) -> Result<GatewayTunnelBaselineReport, Box<dyn std::error::Error>> { ) -> Result<GatewayTunnelBaselineReport, Box<dyn std::error::Error>> {
let tunnel = TunnelHarness::start(TunnelHarnessConfig::default()).await?; let tunnel = TunnelHarness::start(TunnelHarnessConfig {
let peer = connect_protocol_peer(tunnel.base_url()).await?; outbound_queue_capacity: config.outbound_queue_capacity,
..TunnelHarnessConfig::default()
})
.await?;
let mut peers = connect_protocol_peers(tunnel.base_url(), config.tunnel_connections).await?;
let fault_injection = config.close_one_connection_after.map(|delay| {
let peer = peers.pop();
tokio::spawn(async move {
if let Some(peer) = peer {
tokio::time::sleep(delay).await;
peer.abort();
}
})
});
let result = run_http_load_probe(&HttpLoadProbeConfig { let result = run_http_load_probe(&HttpLoadProbeConfig {
url: format!( url: format!(
@@ -90,7 +147,7 @@ async fn run_suite(
"content-type".to_string(), "content-type".to_string(),
"application/octet-stream".to_string(), "application/octet-stream".to_string(),
)]), )]),
body: Some(relay_envelope()), body: Some(relay_envelope(config.request_body_bytes)),
total_requests: config.total_requests, total_requests: config.total_requests,
concurrency: config.concurrency, concurrency: config.concurrency,
timeout: config.timeout, timeout: config.timeout,
@@ -100,16 +157,39 @@ async fn run_suite(
.map_err(std::io::Error::other)?; .map_err(std::io::Error::other)?;
let tunnel_metrics = capture_tunnel_metrics(tunnel.base_url()).await?; let tunnel_metrics = capture_tunnel_metrics(tunnel.base_url()).await?;
drop(peer); if let Some(task) = fault_injection {
let _ = task.await;
}
drop(peers);
let acceptance = evaluate_acceptance(config, &result, &tunnel_metrics);
if config.require_acceptance && !acceptance.passed {
return Err(std::io::Error::other(format!(
"gateway tunnel acceptance failed: {:?}",
acceptance.reasons
))
.into());
}
Ok(GatewayTunnelBaselineReport { Ok(GatewayTunnelBaselineReport {
suite: "gateway_tunnel_stream_baseline", suite: "gateway_tunnel_stream_baseline",
config: GatewayTunnelEffectiveConfig {
total_requests: config.total_requests,
concurrency: config.concurrency,
request_body_bytes: config.request_body_bytes,
tunnel_connections: config.tunnel_connections,
outbound_queue_capacity: config.outbound_queue_capacity,
close_one_connection_after_ms: config
.close_one_connection_after
.map(|duration| duration.as_millis() as u64),
timeout_ms: config.timeout.as_millis() as u64,
},
scenario: result, scenario: result,
tunnel_metrics, tunnel_metrics,
acceptance,
}) })
} }
fn relay_envelope() -> Vec<u8> { fn relay_envelope(request_body_bytes: usize) -> Vec<u8> {
let meta = protocol::RequestMeta { let meta = protocol::RequestMeta {
method: "POST".to_string(), method: "POST".to_string(),
url: "https://baseline.example/v1/chat/completions".to_string(), url: "https://baseline.example/v1/chat/completions".to_string(),
@@ -129,16 +209,29 @@ fn relay_envelope() -> Vec<u8> {
transport_profile: None, transport_profile: None,
}; };
let meta_json = serde_json::to_vec(&meta).expect("tunnel relay metadata should serialize"); let meta_json = serde_json::to_vec(&meta).expect("tunnel relay metadata should serialize");
let body = br#"{"model":"gpt-5","messages":[{"role":"user","content":"hello"}]}"#; let body = vec![b'x'; request_body_bytes];
let mut envelope = Vec::with_capacity(4 + meta_json.len() + body.len()); let mut envelope = Vec::with_capacity(4 + meta_json.len() + body.len());
envelope.extend_from_slice(&(meta_json.len() as u32).to_be_bytes()); envelope.extend_from_slice(&(meta_json.len() as u32).to_be_bytes());
envelope.extend_from_slice(&meta_json); envelope.extend_from_slice(&meta_json);
envelope.extend_from_slice(body); envelope.extend_from_slice(&body);
envelope envelope
} }
async fn connect_protocol_peers(
tunnel_base_url: &str,
count: usize,
) -> Result<Vec<tokio::task::JoinHandle<()>>, Box<dyn std::error::Error>> {
let count = count.max(1);
let mut peers = Vec::with_capacity(count);
for index in 0..count {
peers.push(connect_protocol_peer(tunnel_base_url, index).await?);
}
Ok(peers)
}
async fn connect_protocol_peer( async fn connect_protocol_peer(
tunnel_base_url: &str, tunnel_base_url: &str,
index: usize,
) -> Result<tokio::task::JoinHandle<()>, Box<dyn std::error::Error>> { ) -> Result<tokio::task::JoinHandle<()>, Box<dyn std::error::Error>> {
let ws_url = format!( let ws_url = format!(
"{}{}", "{}{}",
@@ -158,7 +251,7 @@ async fn connect_protocol_peer(
); );
request.headers_mut().insert( request.headers_mut().insert(
"x-node-name", "x-node-name",
http::HeaderValue::from_static("proxy-baseline"), http::HeaderValue::from_str(&format!("proxy-baseline-{index}"))?,
); );
request.headers_mut().insert( request.headers_mut().insert(
"x-tunnel-max-streams", "x-tunnel-max-streams",
@@ -167,6 +260,29 @@ async fn connect_protocol_peer(
let (socket, _response) = tokio_tungstenite::connect_async(request).await?; let (socket, _response) = tokio_tungstenite::connect_async(request).await?;
let (mut sink, mut stream) = socket.split(); let (mut sink, mut stream) = socket.split();
sink.send(Message::Binary(
protocol::encode_hello(&protocol::HelloPayload {
protocol_version: aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION,
capabilities: vec![
"flow-control".to_string(),
"reset-stream".to_string(),
"graceful-drain".to_string(),
],
session_id: Some(format!("baseline-session-{index}")),
replica_id: Some(format!("baseline-replica-{index}")),
})
.into(),
))
.await?;
sink.send(Message::Binary(
protocol::encode_settings(&protocol::SettingsPayload {
initial_stream_window_bytes: 4 * 1024 * 1024,
min_window_update_bytes: 1024 * 1024,
drain_deadline_ms: 30_000,
})
.into(),
))
.await?;
Ok(tokio::spawn(async move { Ok(tokio::spawn(async move {
while let Some(message) = stream.next().await { while let Some(message) = stream.next().await {
let Ok(message) = message else { let Ok(message) = message else {
@@ -250,6 +366,40 @@ async fn capture_tunnel_metrics(
&[], &[],
) )
.unwrap_or_default(), .unwrap_or_default(),
body_backpressure_total: find_metric_value_u64(
&samples,
"tunnel_body_backpressure_total",
&[],
)
.unwrap_or_default(),
flow_window_blocked_ms: find_metric_value_u64(
&samples,
"tunnel_flow_window_blocked_ms",
&[],
)
.unwrap_or_default(),
connection_health_score: find_metric_value_u64(
&samples,
"tunnel_connection_health_score",
&[],
)
.unwrap_or_default(),
stream_reset_total: find_metric_value_u64(
&samples,
"tunnel_stream_reset_total",
&[("reason", "all")],
)
.unwrap_or_default(),
stream_reset_reasons: collect_reason_metrics(&samples, "tunnel_stream_reset_total"),
drain_total: find_metric_value_u64(&samples, "tunnel_drain_total", &[("reason", "all")])
.unwrap_or_default(),
drain_reasons: collect_reason_metrics(&samples, "tunnel_drain_total"),
scheduler_selected_conn_total: find_metric_value_u64(
&samples,
"tunnel_scheduler_selected_conn_total",
&[],
)
.unwrap_or_default(),
proxy_connections_protocol_v1: find_metric_value_u64( proxy_connections_protocol_v1: find_metric_value_u64(
&samples, &samples,
"tunnel_proxy_connections_protocol_v1", "tunnel_proxy_connections_protocol_v1",
@@ -262,9 +412,86 @@ async fn capture_tunnel_metrics(
&[], &[],
) )
.unwrap_or_default(), .unwrap_or_default(),
proxy_connections_protocol_v3: find_metric_value_u64(
&samples,
"tunnel_proxy_connections_protocol_v3",
&[],
)
.unwrap_or_default(),
}) })
} }
fn collect_reason_metrics(
samples: &[PrometheusSample],
metric_name: &str,
) -> BTreeMap<String, u64> {
samples
.iter()
.filter(|sample| {
(sample.name == metric_name || sample.name.ends_with(&format!("_{metric_name}")))
&& sample
.labels
.get("reason")
.is_some_and(|reason| reason != "all")
})
.filter_map(|sample| {
let reason = sample.labels.get("reason")?.clone();
let value = sample.value.parse::<u64>().ok()?;
Some((reason, value))
})
.collect()
}
fn evaluate_acceptance(
config: &GatewayTunnelBaselineConfig,
result: &HttpLoadProbeResult,
metrics: &TunnelMetricsSnapshot,
) -> AcceptanceReport {
let min_success_rate_bps = if config.require_acceptance { 9_950 } else { 1 };
let successful_requests = result
.status_counts
.iter()
.filter(|(status, _)| (200u16..300u16).contains(status))
.map(|(_, count)| *count)
.sum::<usize>();
let success_rate_bps = if result.total_requests == 0 {
0
} else {
((successful_requests as u128) * 10_000 / (result.total_requests as u128)) as u64
};
let congestion_free = metrics.proxy_connection_congested_total == 0;
let no_queue_full_rejections = metrics.outbound_queue_rejected_full_total == 0;
let mut reasons = Vec::new();
if success_rate_bps < min_success_rate_bps {
reasons.push(format!(
"success rate {} bps below required {} bps",
success_rate_bps, min_success_rate_bps
));
}
if !congestion_free {
reasons.push(format!(
"connection congestion total is {}",
metrics.proxy_connection_congested_total
));
}
if !no_queue_full_rejections {
reasons.push(format!(
"outbound queue full rejections total is {}",
metrics.outbound_queue_rejected_full_total
));
}
AcceptanceReport {
required: config.require_acceptance,
passed: reasons.is_empty(),
success_rate_bps,
min_success_rate_bps,
congestion_free,
no_queue_full_rejections,
reasons,
}
}
async fn handle_binary_frame<S>( async fn handle_binary_frame<S>(
sink: &mut S, sink: &mut S,
data: Vec<u8>, data: Vec<u8>,
@@ -285,7 +512,17 @@ where
let payload = protocol::decode_payload(&data, &header).unwrap_or_default(); let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload); let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload);
} }
protocol::REQUEST_BODY if header.flags & protocol::FLAG_END_STREAM != 0 => { protocol::REQUEST_BODY => {
let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
if !payload.is_empty() {
sink.send(Message::Binary(
protocol::encode_window_update(header.stream_id, payload.len() as u32).into(),
))
.await?;
}
if header.flags & protocol::FLAG_END_STREAM == 0 {
return Ok(());
}
let response_meta = protocol::ResponseMeta { let response_meta = protocol::ResponseMeta {
status: 200, status: 200,
headers: vec![( headers: vec![(
@@ -339,6 +576,25 @@ fn parse_args(
"--concurrency" => { "--concurrency" => {
config.concurrency = next_value(&mut iter, "--concurrency")?.parse()? config.concurrency = next_value(&mut iter, "--concurrency")?.parse()?
} }
"--body-bytes" => {
config.request_body_bytes = next_value(&mut iter, "--body-bytes")?.parse()?
}
"--tunnel-connections" => {
config.tunnel_connections =
next_value(&mut iter, "--tunnel-connections")?.parse()?
}
"--outbound-queue-capacity" => {
config.outbound_queue_capacity =
next_value(&mut iter, "--outbound-queue-capacity")?.parse()?
}
"--close-one-connection-after-ms" => {
config.close_one_connection_after = Some(Duration::from_millis(
next_value(&mut iter, "--close-one-connection-after-ms")?.parse()?,
))
}
"--require-acceptance" => {
config.require_acceptance = true;
}
"--timeout-ms" => { "--timeout-ms" => {
config.timeout = config.timeout =
Duration::from_millis(next_value(&mut iter, "--timeout-ms")?.parse()?) Duration::from_millis(next_value(&mut iter, "--timeout-ms")?.parse()?)
@@ -359,6 +615,37 @@ fn parse_args(
} }
} }
} }
if config.request_body_bytes == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"--body-bytes must be positive",
)
.into());
}
if config.tunnel_connections == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"--tunnel-connections must be positive",
)
.into());
}
if config.outbound_queue_capacity == 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"--outbound-queue-capacity must be positive",
)
.into());
}
if config
.close_one_connection_after
.is_some_and(|duration| duration.is_zero())
{
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"--close-one-connection-after-ms must be positive",
)
.into());
}
Ok(config) Ok(config)
} }
@@ -377,6 +664,6 @@ fn next_value(
fn print_usage() { fn print_usage() {
eprintln!( eprintln!(
"usage: cargo run -p aether-testkit --bin gateway_tunnel_stream_baseline -- [--requests 200] [--concurrency 20] [--timeout-ms 10000] [--output /tmp/gateway_tunnel_baseline.json]" "usage: cargo run -p aether-testkit --bin gateway_tunnel_stream_baseline -- [--requests 200] [--concurrency 20] [--body-bytes 6291456] [--tunnel-connections 4] [--outbound-queue-capacity 512] [--close-one-connection-after-ms 1000] [--require-acceptance] [--timeout-ms 10000] [--output /tmp/gateway_tunnel_baseline.json]"
); );
} }
@@ -652,7 +652,30 @@ async fn connect_protocol_peer(
); );
let (socket, _response) = tokio_tungstenite::connect_async(request).await?; let (socket, _response) = tokio_tungstenite::connect_async(request).await?;
let (sink, mut stream) = socket.split(); let (mut sink, mut stream) = socket.split();
sink.send(Message::Binary(
protocol::encode_hello(&protocol::HelloPayload {
protocol_version: aether_contracts::tunnel::CURRENT_TUNNEL_PROTOCOL_VERSION,
capabilities: vec![
"flow-control".to_string(),
"reset-stream".to_string(),
"graceful-drain".to_string(),
],
session_id: Some("llm-stream-stability-session".to_string()),
replica_id: Some("llm-stream-stability-replica".to_string()),
})
.into(),
))
.await?;
sink.send(Message::Binary(
protocol::encode_settings(&protocol::SettingsPayload {
initial_stream_window_bytes: 4 * 1024 * 1024,
min_window_update_bytes: 1024 * 1024,
drain_deadline_ms: 30_000,
})
.into(),
))
.await?;
let sink = Arc::new(Mutex::new(sink)); let sink = Arc::new(Mutex::new(sink));
let stats = Arc::new(PeerStats::default()); let stats = Arc::new(PeerStats::default());
let cancelled = Arc::new(Mutex::new(HashSet::new())); let cancelled = Arc::new(Mutex::new(HashSet::new()));
@@ -745,7 +768,21 @@ async fn handle_peer_binary_frame(
let payload = protocol::decode_payload(&data, &header).unwrap_or_default(); let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload); let _ = serde_json::from_slice::<protocol::RequestMeta>(&payload);
} }
protocol::REQUEST_BODY if header.flags & protocol::FLAG_END_STREAM != 0 => { protocol::REQUEST_BODY => {
let payload = protocol::decode_payload(&data, &header).unwrap_or_default();
if !payload.is_empty() {
let _ = send_ws_message(
&sink,
Message::Binary(
protocol::encode_window_update(header.stream_id, payload.len() as u32)
.into(),
),
)
.await;
}
if header.flags & protocol::FLAG_END_STREAM == 0 {
return;
}
stats stats
.request_body_end_received .request_body_end_received
.fetch_add(1, Ordering::AcqRel); .fetch_add(1, Ordering::AcqRel);
@@ -767,6 +804,10 @@ async fn handle_peer_binary_frame(
stats.stream_errors_received.fetch_add(1, Ordering::AcqRel); stats.stream_errors_received.fetch_add(1, Ordering::AcqRel);
cancelled.lock().await.insert(header.stream_id); cancelled.lock().await.insert(header.stream_id);
} }
protocol::RESET_STREAM => {
stats.stream_errors_received.fetch_add(1, Ordering::AcqRel);
cancelled.lock().await.insert(header.stream_id);
}
_ => {} _ => {}
} }
} }