refactor: 移除独立 hub/proxy/executor/gateway crate,统一为 gateway tunnel 架构

- 删除 aether-hub、aether-proxy 独立项目及其 Dockerfile/配置
- 删除 crates/aether-executor 和 crates/aether-gateway 全部模块
- 新增 apps/ 目录作为应用入口
- 将 hub 概念重构为 gateway tunnel transport
- 将 executor 重构为 execution runtime
- 新增 tunnel.rs 合约定义和 testkit tunnel/execution_runtime 模块
- 更新 Python 服务层和测试适配新架构命名
This commit is contained in:
fawney19
2026-04-03 14:59:58 +08:00
parent ddf18fed9a
commit 8f26e1a31f
983 changed files with 103097 additions and 105836 deletions
@@ -0,0 +1,141 @@
use aether_http::{build_http_client, HttpClientConfig};
use futures_util::future::BoxFuture;
use reqwest::Client;
use std::sync::Arc;
type HeartbeatAckCallback =
dyn Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync;
type NodeStatusCallback =
dyn Fn(String, bool, usize) -> BoxFuture<'static, Result<(), String>> + Send + Sync;
enum ControlPlaneMode {
Disabled,
Http {
client: Option<Client>,
base_url: String,
},
Local {
heartbeat_ack: Arc<HeartbeatAckCallback>,
push_node_status: Arc<NodeStatusCallback>,
},
}
#[derive(Clone)]
pub struct ControlPlaneClient {
inner: Arc<ControlPlaneMode>,
}
impl ControlPlaneClient {
pub fn new(base_url: String) -> Self {
let client = build_http_client(&HttpClientConfig {
request_timeout_ms: Some(10_000),
user_agent: Some("aether-tunnel-standalone/control-plane".to_string()),
..HttpClientConfig::default()
})
.ok();
Self {
inner: Arc::new(ControlPlaneMode::Http { client, base_url }),
}
}
pub fn disabled() -> Self {
Self {
inner: Arc::new(ControlPlaneMode::Disabled),
}
}
pub fn local<HeartbeatAck, PushNodeStatus>(
heartbeat_ack: HeartbeatAck,
push_node_status: PushNodeStatus,
) -> Self
where
HeartbeatAck:
Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync + 'static,
PushNodeStatus: Fn(String, bool, usize) -> BoxFuture<'static, Result<(), String>>
+ Send
+ Sync
+ 'static,
{
Self {
inner: Arc::new(ControlPlaneMode::Local {
heartbeat_ack: Arc::new(heartbeat_ack),
push_node_status: Arc::new(push_node_status),
}),
}
}
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
match self.inner.as_ref() {
ControlPlaneMode::Disabled => Ok(b"{}".to_vec()),
ControlPlaneMode::Http { client, base_url } => {
let Some(client) = client else {
return Ok(b"{}".to_vec());
};
let url = format!(
"{}/api/internal/tunnel/heartbeat",
base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.header("content-type", "application/json")
.body(payload.to_vec())
.send()
.await
.map_err(|e| format!("heartbeat callback request failed: {e}"))?;
if !response.status().is_success() {
return Err(format!(
"heartbeat callback failed with status {}",
response.status()
));
}
response
.bytes()
.await
.map(|bytes| bytes.to_vec())
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
}
ControlPlaneMode::Local { heartbeat_ack, .. } => heartbeat_ack(payload.to_vec()).await,
}
}
pub async fn push_node_status(
&self,
node_id: &str,
connected: bool,
conn_count: usize,
) -> Result<(), String> {
match self.inner.as_ref() {
ControlPlaneMode::Disabled => Ok(()),
ControlPlaneMode::Http { client, base_url } => {
let Some(client) = client else {
return Ok(());
};
let url = format!(
"{}/api/internal/tunnel/node-status",
base_url.trim_end_matches('/')
);
let response = client
.post(&url)
.json(&serde_json::json!({
"node_id": node_id,
"connected": connected,
"conn_count": conn_count,
}))
.send()
.await
.map_err(|e| format!("node-status callback request failed: {e}"))?;
if response.status().is_success() {
Ok(())
} else {
Err(format!(
"node-status callback failed with status {}",
response.status()
))
}
}
ControlPlaneMode::Local {
push_node_status, ..
} => push_node_status(node_id.to_string(), connected, conn_count).await,
}
}
}
@@ -0,0 +1,922 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use aether_runtime::{BoundedQueueSender, MetricKind, MetricSample, QueueSendError};
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};
use super::control_plane::ControlPlaneClient;
use super::protocol;
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SendStatus {
Queued,
Closed,
Congested,
}
#[derive(Debug, Clone, Copy)]
pub struct ConnConfig {
pub ping_interval: Duration,
pub idle_timeout: Duration,
pub outbound_queue_capacity: usize,
}
pub struct BoundedOutbound {
tx: BoundedQueueSender<Message>,
close_tx: watch::Sender<bool>,
closing: AtomicBool,
}
impl BoundedOutbound {
pub fn new(tx: BoundedQueueSender<Message>, close_tx: watch::Sender<bool>) -> Self {
Self {
tx,
close_tx,
closing: AtomicBool::new(false),
}
}
pub fn send(&self, msg: Message) -> SendStatus {
if self.is_closing() {
return SendStatus::Closed;
}
match self.tx.try_send(msg) {
Ok(()) => SendStatus::Queued,
Err(QueueSendError::Closed(_)) => {
self.mark_closing();
SendStatus::Closed
}
Err(QueueSendError::Full(_)) => {
self.mark_closing();
SendStatus::Congested
}
}
}
pub fn is_closing(&self) -> bool {
self.closing.load(Ordering::Acquire)
}
pub fn mark_closing(&self) -> bool {
if self.closing.swap(true, Ordering::AcqRel) {
return false;
}
let _ = self.close_tx.send(true);
true
}
}
pub struct ProxyConn {
pub id: u64,
pub node_id: String,
pub node_name: String,
pub outbound: BoundedOutbound,
next_stream_id: AtomicU32,
pub stream_count: AtomicUsize,
pub max_streams: usize,
}
impl ProxyConn {
pub fn new(
id: u64,
node_id: String,
node_name: String,
tx: BoundedQueueSender<Message>,
close_tx: watch::Sender<bool>,
max_streams: usize,
) -> Self {
Self {
id,
node_id,
node_name,
outbound: BoundedOutbound::new(tx, close_tx),
next_stream_id: AtomicU32::new(2),
stream_count: AtomicUsize::new(0),
max_streams,
}
}
pub fn alloc_stream_id(&self) -> Option<u32> {
let mut current = self.stream_count.load(Ordering::Relaxed);
loop {
if current >= self.max_streams || !self.is_available() {
return None;
}
match self.stream_count.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
let sid = loop {
let current_sid = self.next_stream_id.load(Ordering::Relaxed);
let next_sid = if current_sid >= 0xFFFF_FFFE {
2
} else {
current_sid + 2
};
if self
.next_stream_id
.compare_exchange_weak(current_sid, next_sid, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
break current_sid;
}
};
Some(sid)
}
pub fn release_stream(&self) {
let mut current = self.stream_count.load(Ordering::Relaxed);
while current > 0 {
match self.stream_count.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
pub fn is_available(&self) -> bool {
!self.outbound.is_closing()
}
pub fn request_close(&self) {
self.outbound.mark_closing();
}
pub fn send(&self, msg: Message) -> SendStatus {
let was_closing = self.outbound.is_closing();
let status = self.outbound.send(msg);
if status == SendStatus::Congested && !was_closing {
warn!(
conn_id = self.id,
node_id = %self.node_id,
node_name = %self.node_name,
queued_streams = self.stream_count.load(Ordering::Relaxed),
"proxy outbound queue full, closing congested connection"
);
}
status
}
}
#[derive(Debug, Clone)]
pub struct LocalResponseHead {
pub status: u16,
pub headers: Vec<(String, String)>,
}
#[derive(Debug)]
pub enum LocalBodyEvent {
Chunk(Bytes),
End,
Error(String),
}
#[derive(Debug, Default)]
struct LocalWaitState {
response: Option<LocalResponseHead>,
error: Option<String>,
}
pub struct LocalStream {
pub id: u64,
proxy_conn_id: u64,
proxy_stream_id: u32,
wait_state: Mutex<LocalWaitState>,
headers_notify: Notify,
body_tx: mpsc::Sender<LocalBodyEvent>,
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
terminal: AtomicBool,
}
impl LocalStream {
fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32) -> Self {
let (body_tx, body_rx) = mpsc::channel(128);
Self {
id,
proxy_conn_id,
proxy_stream_id,
wait_state: Mutex::new(LocalWaitState::default()),
headers_notify: Notify::new(),
body_tx,
body_rx: Mutex::new(Some(body_rx)),
terminal: AtomicBool::new(false),
}
}
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
tokio::time::timeout(timeout, async {
loop {
let outcome = {
let state = self.wait_state.lock();
if let Some(response) = &state.response {
return Ok(response.clone());
}
state.error.clone()
};
if let Some(error) = outcome {
return Err(error);
}
self.headers_notify.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()
}
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.response = Some(LocalResponseHead {
status: meta.status,
headers: meta.headers,
});
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
}
fn push_body_chunk(&self, payload: Bytes) -> bool {
if self.terminal.load(Ordering::Acquire) {
return false;
}
self.body_tx
.try_send(LocalBodyEvent::Chunk(payload))
.is_ok()
}
fn finish(&self) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some("stream ended before response headers".to_string());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::End);
}
fn fail(&self, error: impl Into<String>) {
if self.terminal.swap(true, Ordering::AcqRel) {
return;
}
let error = error.into();
let mut notify = false;
{
let mut state = self.wait_state.lock();
if state.response.is_none() && state.error.is_none() {
state.error = Some(error.clone());
notify = true;
}
}
if notify {
self.headers_notify.notify_waiters();
}
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
}
}
pub struct HubRouter {
proxy_conns: RwLock<HashMap<String, Vec<Arc<ProxyConn>>>>,
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
local_streams: DashMap<u64, Arc<LocalStream>>,
proxy_to_local: DashMap<(u64, u32), u64>,
next_conn_id: AtomicU64,
next_local_stream_id: AtomicU64,
control_plane: ControlPlaneClient,
}
impl HubRouter {
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
Arc::new(Self {
proxy_conns: RwLock::new(HashMap::new()),
proxy_conns_by_id: DashMap::new(),
local_streams: DashMap::new(),
proxy_to_local: DashMap::new(),
next_conn_id: AtomicU64::new(1),
next_local_stream_id: AtomicU64::new(1),
control_plane,
})
}
pub fn alloc_conn_id(&self) -> u64 {
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
}
pub fn register_proxy(&self, conn: Arc<ProxyConn>) {
let node_id = conn.node_id.clone();
let node_name = conn.node_name.clone();
let conn_id = conn.id;
self.proxy_conns_by_id.insert(conn_id, conn.clone());
let pool_size = {
let mut map = self.proxy_conns.write();
map.entry(node_id.clone()).or_default().push(conn);
map.get(&node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
node_name = %node_name,
conn_id = conn_id,
pool_size = pool_size,
"proxy connected"
);
self.notify_node_status(node_id, true, pool_size);
}
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
self.proxy_conns_by_id.remove(&conn_id);
let pool_size = {
let mut map = self.proxy_conns.write();
if let Some(conns) = map.get_mut(node_id) {
conns.retain(|c| c.id != conn_id);
if conns.is_empty() {
map.remove(node_id);
}
}
map.get(node_id).map(|v| v.len()).unwrap_or(0)
};
info!(
node_id = %node_id,
conn_id = conn_id,
remaining = pool_size,
"proxy disconnected"
);
self.cancel_streams_for_proxy(conn_id);
self.notify_node_status(node_id.to_string(), pool_size > 0, pool_size);
}
pub fn request_close_all_proxies(&self) -> usize {
let conns = self
.proxy_conns_by_id
.iter()
.map(|entry| Arc::clone(entry.value()))
.collect::<Vec<_>>();
let total = conns.len();
for conn in conns {
conn.request_close();
}
total
}
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
let control_plane = self.control_plane.clone();
tokio::spawn(async move {
if let Err(error) = control_plane
.push_node_status(&node_id, connected, conn_count)
.await
{
warn!(
node_id = %node_id,
connected = connected,
conn_count = conn_count,
error = %error,
"failed to push node status to app control plane"
);
}
});
}
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
let map = self.proxy_conns.read();
let conns = map.get(node_id)?;
conns
.iter()
.filter(|c| c.is_available())
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned()
}
pub fn has_local_proxy(&self, node_id: &str) -> bool {
self.get_proxy_conn(node_id).is_some()
}
pub fn open_local_stream(
&self,
node_id: &str,
meta: &protocol::RequestMeta,
) -> Result<Arc<LocalStream>, String> {
let proxy_conn = self
.get_proxy_conn(node_id)
.ok_or_else(|| format!("no proxy connection for node {node_id}"))?;
let proxy_stream_id = proxy_conn
.alloc_stream_id()
.ok_or_else(|| format!("stream limit reached for node {node_id}"))?;
// Encode frames before registering the stream so that encoding failures
// (practically impossible but theoretically possible) don't leak a stream
// slot or orphan map entries.
let meta_json = match serde_json::to_vec(meta) {
Ok(json) => json,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to encode request metadata: {e}"));
}
};
let (meta_payload, meta_flags) = match protocol::compress_payload(&meta_json) {
Ok(result) => result,
Err(e) => {
proxy_conn.release_stream();
return Err(format!("failed to compress request metadata: {e}"));
}
};
let header_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_HEADERS,
meta_flags,
&meta_payload,
);
// 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(
local_stream_id,
proxy_conn.id,
proxy_stream_id,
));
self.local_streams
.insert(local_stream_id, local_stream.clone());
self.proxy_to_local
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
match proxy_conn.send(Message::Binary(header_frame.into())) {
SendStatus::Queued => Ok(local_stream),
SendStatus::Closed | SendStatus::Congested => {
self.cleanup_local_stream(local_stream_id);
proxy_conn.release_stream();
Err("proxy connection congested".to_string())
}
}
}
pub fn push_local_request_body(
&self,
local_stream_id: u64,
payload: Bytes,
end_stream: bool,
) -> Result<(), String> {
let stream = self
.local_streams
.get(&local_stream_id)
.map(|entry| entry.value().clone())
.ok_or_else(|| "local stream not found".to_string())?;
let proxy_conn = self
.proxy_conns_by_id
.get(&stream.proxy_conn_id)
.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);
if total_chunks == 0 {
if end_stream {
self.send_request_body_frame(&proxy_conn, stream.proxy_stream_id, &[], true)?;
}
} else {
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
let is_last_chunk = index + 1 == total_chunks;
self.send_request_body_frame(
&proxy_conn,
stream.proxy_stream_id,
chunk,
end_stream && is_last_chunk,
)?;
}
}
Ok(())
}
fn send_request_body_frame(
&self,
proxy_conn: &Arc<ProxyConn>,
proxy_stream_id: u32,
payload: &[u8],
end_stream: bool,
) -> Result<(), String> {
let (body_payload, body_flags) = protocol::compress_payload(payload)
.map_err(|e| format!("failed to compress request body: {e}"))?;
let body_frame = protocol::encode_frame(
proxy_stream_id,
protocol::REQUEST_BODY,
body_flags
| if end_stream {
protocol::FLAG_END_STREAM
} else {
0
},
&body_payload,
);
match proxy_conn.send(Message::Binary(body_frame.into())) {
SendStatus::Queued => Ok(()),
SendStatus::Closed | SendStatus::Congested => {
Err("proxy connection congested".to_string())
}
}
}
pub fn cancel_local_stream(&self, local_stream_id: u64, reason: &str) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
if let Some(pc) = self.proxy_conns_by_id.get(&stream.proxy_conn_id) {
pc.release_stream();
let frame = protocol::encode_stream_error(stream.proxy_stream_id, reason);
let _ = pc.send(Message::Binary(frame.into()));
}
stream.fail(reason.to_string());
}
fn cleanup_local_stream(&self, local_stream_id: u64) {
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
return;
};
self.proxy_to_local
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
}
pub async fn handle_proxy_frame(&self, proxy_conn_id: u64, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return;
}
match header.msg_type {
protocol::RESPONSE_HEADERS => {
self.route_response_headers(proxy_conn_id, header, data);
}
protocol::RESPONSE_BODY => {
self.route_response_body(proxy_conn_id, header, data);
}
protocol::STREAM_END => {
self.finish_proxy_stream(proxy_conn_id, header.stream_id);
}
protocol::STREAM_ERROR => {
let message = protocol::decode_payload(data, &header)
.ok()
.and_then(|payload| String::from_utf8(payload).ok())
.unwrap_or_else(|| "stream error".to_string());
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
}
protocol::HEARTBEAT_DATA => {
self.handle_heartbeat(proxy_conn_id, header.stream_id, data, &header)
.await;
}
protocol::PING => {
let payload = protocol::frame_payload_by_header(data, &header).unwrap_or(&[]);
let pong = protocol::encode_pong(payload);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(pong.into()));
}
}
protocol::PONG => {}
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
_ => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"unexpected frame type from proxy"
);
}
}
}
fn route_response_headers(
&self,
proxy_conn_id: u64,
header: protocol::FrameHeader,
data: &[u8],
) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response headers",
);
return;
};
let Ok(meta) = serde_json::from_slice::<protocol::ResponseMeta>(&payload) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"invalid response headers payload",
);
return;
};
if let Some(entry) = self.local_streams.get(&local_id) {
entry.value().set_response_headers(meta);
}
}
fn route_response_body(&self, proxy_conn_id: u64, header: protocol::FrameHeader, data: &[u8]) {
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
return;
};
let Ok(payload) = protocol::decode_payload(data, &header) else {
self.fail_proxy_stream(
proxy_conn_id,
header.stream_id,
"failed to decode response body",
);
return;
};
let stream = match self.local_streams.get(&local_id) {
Some(entry) => entry.value().clone(),
None => return,
};
if !stream.push_body_chunk(Bytes::from(payload)) {
self.cancel_local_stream(local_id, "local relay response congested");
}
}
fn handle_stream_cleanup(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
) -> Option<Arc<LocalStream>> {
let local_id = self
.proxy_to_local
.remove(&(proxy_conn_id, proxy_stream_id))
.map(|(_, local_id)| local_id)?;
let stream = self
.local_streams
.remove(&local_id)
.map(|(_, stream)| stream)?;
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
pc.release_stream();
}
Some(stream)
}
fn finish_proxy_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.finish();
}
}
fn fail_proxy_stream(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
error: impl Into<String>,
) {
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
stream.fail(error.into());
}
}
fn lookup_local_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) -> Option<u64> {
self.proxy_to_local
.get(&(proxy_conn_id, proxy_stream_id))
.map(|entry| *entry.value())
}
async fn handle_heartbeat(
&self,
proxy_conn_id: u64,
stream_id: u32,
data: &[u8],
header: &protocol::FrameHeader,
) {
let payload = match protocol::decode_payload(data, header) {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "failed to decode heartbeat payload");
return;
}
};
let ack_payload = match self.control_plane.heartbeat_ack(&payload).await {
Ok(payload) => payload,
Err(error) => {
warn!(proxy_conn_id = proxy_conn_id, error = %error, "control-plane heartbeat callback failed");
b"{}".to_vec()
}
};
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let frame = protocol::encode_frame(stream_id, protocol::HEARTBEAT_ACK, 0, &ack_payload);
let _ = pc.send(Message::Binary(frame.into()));
}
}
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
let mut cancelled = 0usize;
self.proxy_to_local.retain(|key, local_id| {
if key.0 != proxy_conn_id {
return true;
}
if let Some((_, stream)) = self.local_streams.remove(local_id) {
stream.fail("proxy disconnected".to_string());
}
cancelled += 1;
false
});
if cancelled > 0 {
warn!(
proxy_conn_id = proxy_conn_id,
streams_cancelled = cancelled,
"cancelled in-flight streams due to proxy disconnect"
);
}
}
pub fn stats(&self) -> HubStats {
let proxy_conns = self.proxy_conns.read();
let total_proxy = proxy_conns.values().map(|v| v.len()).sum();
let nodes = proxy_conns.len();
drop(proxy_conns);
HubStats {
proxy_connections: total_proxy,
nodes,
active_streams: self.local_streams.len(),
}
}
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub nodes: usize,
pub active_streams: usize,
}
impl HubStats {
pub fn to_metric_samples(&self) -> Vec<MetricSample> {
vec![
MetricSample::new(
"tunnel_proxy_connections",
"Current number of connected proxy sockets.",
MetricKind::Gauge,
self.proxy_connections as u64,
),
MetricSample::new(
"tunnel_nodes",
"Current number of connected logical nodes.",
MetricKind::Gauge,
self.nodes as u64,
),
MetricSample::new(
"tunnel_active_streams",
"Current number of active local relay streams.",
MetricKind::Gauge,
self.active_streams as u64,
),
]
}
}
#[cfg(test)]
mod tests {
use aether_runtime::bounded_queue;
use super::*;
fn build_meta() -> protocol::RequestMeta {
protocol::RequestMeta {
method: "GET".to_string(),
url: "https://example.com".to_string(),
headers: HashMap::new(),
timeout: 30,
}
}
#[tokio::test]
async fn cancel_local_stream_notifies_proxy() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
100,
"node-1".to_string(),
"Node 1".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-1", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
hub.push_local_request_body(stream.id, Bytes::new(), true)
.expect("finish empty body");
let _ = proxy_rx.try_recv().expect("body frame");
hub.cancel_local_stream(stream.id, "client dropped");
let cancelled = proxy_rx.try_recv().expect("cancel frame");
let cancelled_data = match cancelled {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
}
#[tokio::test]
async fn push_local_request_body_splits_large_payload_and_marks_end() {
let hub = HubRouter::new(ControlPlaneClient::disabled());
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
let (proxy_close_tx, _) = watch::channel(false);
let proxy = Arc::new(ProxyConn::new(
200,
"node-2".to_string(),
"Node 2".to_string(),
proxy_tx,
proxy_close_tx,
16,
));
hub.register_proxy(proxy);
let stream = hub
.open_local_stream("node-2", &build_meta())
.expect("open local stream");
let _ = proxy_rx.try_recv().expect("headers frame");
let payload = Bytes::from(vec![b'x'; MAX_REQUEST_BODY_FRAME_SIZE + 17]);
hub.push_local_request_body(stream.id, payload, true)
.expect("push request body");
let first = match proxy_rx.try_recv().expect("first body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let first_header = protocol::FrameHeader::parse(&first).expect("first body header");
assert_eq!(first_header.msg_type, protocol::REQUEST_BODY);
assert_eq!(first_header.flags & protocol::FLAG_END_STREAM, 0);
let second = match proxy_rx.try_recv().expect("second body frame") {
Message::Binary(data) => data.to_vec(),
other => panic!("unexpected message: {other:?}"),
};
let second_header = protocol::FrameHeader::parse(&second).expect("second body header");
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
}
}
@@ -0,0 +1,398 @@
use std::io;
use std::net::SocketAddr;
use std::time::Duration;
use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER;
use aether_runtime::{maybe_hold_axum_response_permit, AdmissionPermit};
use async_stream::stream;
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 bytes::BytesMut;
use futures_util::StreamExt;
use tracing::warn;
use super::hub::{LocalBodyEvent, LocalStream};
use super::protocol;
use super::AppState;
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
const MAX_RELAY_META_LEN: usize = 256 * 1024;
struct StreamGuard {
hub: std::sync::Arc<super::hub::HubRouter>,
stream_id: u64,
finished: bool,
}
impl Drop for StreamGuard {
fn drop(&mut self) {
if !self.finished {
self.hub
.cancel_local_stream(self.stream_id, "local relay client dropped");
}
}
}
pub async fn relay_request(
Path(node_id): Path<String>,
State(state): State<AppState>,
ConnectInfo(addr): ConnectInfo<SocketAddr>,
request: Request,
) -> impl IntoResponse {
let forwarded_by_gateway = request
.headers()
.get(TUNNEL_RELAY_FORWARDED_BY_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if !addr.ip().is_loopback() && !forwarded_by_gateway {
return tunnel_error_response(
StatusCode::FORBIDDEN,
"forbidden",
"local relay only accepts loopback requests",
);
}
let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
Err(super::RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
..
}))
| Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Saturated { .. },
))
| Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Unavailable { .. },
)) => {
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"overloaded",
"hub relay overloaded",
);
}
Err(super::RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed {
..
})) => {
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"overloaded",
"hub relay gate closed",
);
}
Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::InvalidConfiguration(_),
)) => {
return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE,
"overloaded",
"hub relay distributed gate invalid",
);
}
};
let mut body_stream = request.into_body().into_data_stream();
let mut envelope_buf = BytesMut::new();
let mut meta: Option<protocol::RequestMeta> = None;
let mut stream: Option<std::sync::Arc<LocalStream>> = None;
while let Some(chunk_result) = body_stream.next().await {
let chunk = match chunk_result {
Ok(chunk) => chunk,
Err(error) => {
if let Some(active_stream) = &stream {
state
.hub
.cancel_local_stream(active_stream.id, "failed to read relay request body");
}
warn!(error = %error, "failed to read local relay request body");
return release_permit_response(
tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to read relay request body",
),
request_permit,
);
}
};
if stream.is_none() {
envelope_buf.extend_from_slice(&chunk);
let Some((parsed_meta, body_offset)) = (match try_decode_envelope_meta(&envelope_buf) {
Ok(result) => result,
Err(error) => {
return release_permit_response(
tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error),
request_permit,
);
}
}) else {
continue;
};
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) {
Ok(stream) => stream,
Err(error) => {
return release_permit_response(
tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error),
request_permit,
);
}
};
if envelope_buf.len() > body_offset {
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
if let Err(error) =
state
.hub
.push_local_request_body(opened_stream.id, first_body_chunk, false)
{
state.hub.cancel_local_stream(opened_stream.id, &error);
return release_permit_response(
tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error),
request_permit,
);
}
}
envelope_buf.clear();
meta = Some(parsed_meta);
stream = Some(opened_stream);
continue;
}
let Some(active_stream) = &stream else {
continue;
};
if let Err(error) = state
.hub
.push_local_request_body(active_stream.id, chunk, false)
{
state.hub.cancel_local_stream(active_stream.id, &error);
return release_permit_response(
tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error),
request_permit,
);
}
}
let (meta, stream) = match (meta, stream) {
(Some(meta), Some(stream)) => (meta, stream),
_ => {
return release_permit_response(
tunnel_error_response(
StatusCode::BAD_REQUEST,
"bad_request",
"relay envelope metadata truncated",
),
request_permit,
);
}
};
if let Err(error) = state
.hub
.push_local_request_body(stream.id, Bytes::new(), true)
{
state.hub.cancel_local_stream(stream.id, &error);
return release_permit_response(
tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error),
request_permit,
);
}
let request_guard = StreamGuard {
hub: state.hub.clone(),
stream_id: stream.id,
finished: false,
};
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
let response_head = match stream.wait_headers(wait_timeout).await {
Ok(response) => response,
Err(error) => {
state.hub.cancel_local_stream(stream.id, &error);
return release_permit_response(
tunnel_error_response(StatusCode::GATEWAY_TIMEOUT, "timeout", &error),
request_permit,
);
}
};
let Some(mut body_rx) = stream.take_body_receiver() else {
state
.hub
.cancel_local_stream(stream.id, "missing relay response body receiver");
return release_permit_response(
tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"missing relay response body receiver",
),
request_permit,
);
};
let hub = state.hub.clone();
let stream_id = stream.id;
let body_stream = stream! {
let mut guard = request_guard;
guard.hub = hub;
guard.stream_id = stream_id;
while let Some(event) = body_rx.recv().await {
match event {
LocalBodyEvent::Chunk(chunk) => yield Ok::<Bytes, io::Error>(chunk),
LocalBodyEvent::End => {
guard.finished = true;
break;
}
LocalBodyEvent::Error(error) => {
guard.finished = true;
yield Err(io::Error::other(error));
break;
}
}
}
guard.finished = true;
};
let mut builder = Response::builder().status(response_head.status);
if let Some(headers) = builder.headers_mut() {
append_headers(headers, &response_head.headers);
}
match builder.body(Body::from_stream(body_stream)) {
Ok(response) => maybe_hold_axum_response_permit(response, request_permit),
Err(error) => {
warn!(error = %error, "failed to build relay response");
release_permit_response(
tunnel_error_response(
StatusCode::BAD_GATEWAY,
"relay",
"failed to build relay response",
),
request_permit,
)
}
}
}
fn release_permit_response(
response: Response<Body>,
_request_permit: Option<AdmissionPermit>,
) -> Response<Body> {
response
}
fn try_decode_envelope_meta(
buffer: &BytesMut,
) -> Result<Option<(protocol::RequestMeta, usize)>, String> {
if buffer.len() < 4 {
return Ok(None);
}
let meta_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
if meta_len > MAX_RELAY_META_LEN {
return Err("relay metadata too large".to_string());
}
let meta_end = 4usize
.checked_add(meta_len)
.ok_or_else(|| "relay envelope length overflow".to_string())?;
if buffer.len() < meta_end {
return Ok(None);
}
let meta = serde_json::from_slice::<protocol::RequestMeta>(&buffer[4..meta_end])
.map_err(|e| format!("invalid relay metadata: {e}"))?;
Ok(Some((meta, meta_end)))
}
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
for (name, value) in headers {
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
continue;
};
let Ok(value) = HeaderValue::from_str(value) else {
continue;
};
target.append(name, value);
}
}
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
let mut builder = Response::builder().status(status);
if let Some(headers) = builder.headers_mut() {
headers.insert(
HeaderName::from_static(TUNNEL_ERROR_HEADER),
HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")),
);
headers.insert(
axum::http::header::CONTENT_TYPE,
HeaderValue::from_static("text/plain; charset=utf-8"),
);
}
builder
.body(Body::from(message.to_string()))
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
}
#[cfg(test)]
mod tests {
use super::super::{AppState, ConnConfig, ControlPlaneClient};
use super::*;
use axum::extract::{ConnectInfo, Path, State};
use axum::response::IntoResponse;
fn test_app_state() -> AppState {
AppState::new(
ControlPlaneClient::disabled(),
ConnConfig {
ping_interval: Duration::from_secs(15),
idle_timeout: Duration::from_secs(0),
outbound_queue_capacity: 128,
},
128,
)
}
#[tokio::test]
async fn relay_rejects_non_loopback_without_forwarded_header() {
let request = Request::builder()
.body(Body::empty())
.expect("request should build");
let response = relay_request(
Path("node-123".to_string()),
State(test_app_state()),
ConnectInfo(SocketAddr::from(([10, 0, 0, 1], 4242))),
request,
)
.await
.into_response();
assert_eq!(response.status(), StatusCode::FORBIDDEN);
}
#[tokio::test]
async fn relay_accepts_forwarded_gateway_request_from_non_loopback() {
let request = Request::builder()
.header(TUNNEL_RELAY_FORWARDED_BY_HEADER, "gateway-a")
.body(Body::empty())
.expect("request should build");
let response = relay_request(
Path("node-123".to_string()),
State(test_app_state()),
ConnectInfo(SocketAddr::from(([10, 0, 0, 1], 4242))),
request,
)
.await
.into_response();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
assert_eq!(
response
.headers()
.get(TUNNEL_ERROR_HEADER)
.and_then(|value| value.to_str().ok()),
Some("bad_request")
);
}
}
@@ -0,0 +1,255 @@
mod control_plane;
mod hub;
mod local_relay;
pub mod protocol;
mod proxy_conn;
use std::sync::Arc;
use aether_runtime::{
hold_admission_permit_until, prometheus_response, service_up_sample, AdmissionPermit,
ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, DistributedConcurrencyError,
DistributedConcurrencyGate, DistributedConcurrencySnapshot, MetricKind, MetricLabel,
MetricSample,
};
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State;
use axum::response::{IntoResponse, Json};
use axum::routing::{get, post};
use axum::Router;
use tracing::warn;
pub use control_plane::ControlPlaneClient;
pub use hub::{ConnConfig, HubRouter};
pub use local_relay::relay_request;
#[derive(Clone)]
pub struct AppState {
pub hub: Arc<HubRouter>,
pub proxy_conn_cfg: ConnConfig,
pub max_streams: usize,
request_gate: Option<Arc<ConcurrencyGate>>,
distributed_request_gate: Option<Arc<DistributedConcurrencyGate>>,
}
#[derive(Debug)]
enum RequestAdmissionError {
Local(ConcurrencyError),
Distributed(DistributedConcurrencyError),
}
impl AppState {
pub fn new(
control_plane: ControlPlaneClient,
proxy_conn_cfg: ConnConfig,
max_streams: usize,
) -> Self {
Self {
hub: HubRouter::new(control_plane),
proxy_conn_cfg,
max_streams,
request_gate: None,
distributed_request_gate: None,
}
}
pub fn with_request_concurrency_limit(mut self, limit: Option<usize>) -> Self {
self.request_gate = limit
.filter(|limit| *limit > 0)
.map(|limit| Arc::new(ConcurrencyGate::new("tunnel_requests", limit)));
self
}
pub fn with_distributed_request_gate(mut self, gate: DistributedConcurrencyGate) -> Self {
self.distributed_request_gate = Some(Arc::new(gate));
self
}
fn request_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.request_gate.as_ref().map(|gate| gate.snapshot())
}
async fn distributed_request_concurrency_snapshot(
&self,
) -> Result<Option<DistributedConcurrencySnapshot>, DistributedConcurrencyError> {
match self.distributed_request_gate.as_ref() {
Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None),
}
}
async fn metric_samples(&self) -> Vec<MetricSample> {
let mut samples = vec![service_up_sample("aether-tunnel-standalone")];
if let Some(snapshot) = self.request_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("tunnel_requests"));
}
if let Some(gate) = self.distributed_request_gate.as_ref() {
match gate.snapshot().await {
Ok(snapshot) => {
samples.extend(snapshot.to_metric_samples("tunnel_requests_distributed"));
}
Err(_) => samples.push(
MetricSample::new(
"concurrency_unavailable",
"Whether the distributed concurrency gate is currently unavailable.",
MetricKind::Gauge,
1,
)
.with_labels(vec![MetricLabel::new(
"gate",
"tunnel_requests_distributed",
)]),
),
}
}
samples.extend(self.hub.stats().to_metric_samples());
samples
}
async fn try_acquire_request_permit(
&self,
) -> Result<Option<AdmissionPermit>, RequestAdmissionError> {
let local = self
.request_gate
.as_ref()
.map(|gate| gate.try_acquire())
.transpose()
.map_err(RequestAdmissionError::Local)?;
let distributed = match self.distributed_request_gate.as_ref() {
Some(gate) => Some(
gate.try_acquire()
.await
.map_err(RequestAdmissionError::Distributed)?,
),
None => None,
};
Ok(AdmissionPermit::from_parts(local, distributed))
}
}
pub fn build_router_with_state(state: AppState) -> Router {
Router::new()
.route("/health", get(health))
.route("/metrics", get(metrics))
.route("/stats", get(stats))
.route("/api/internal/proxy-tunnel", get(ws_proxy))
.route(
"/api/internal/tunnel/relay/{node_id}",
post(local_relay::relay_request),
)
.with_state(state)
}
async fn health(State(state): State<AppState>) -> impl IntoResponse {
let request_concurrency = state.request_concurrency_snapshot().map(|snapshot| {
serde_json::json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
let distributed_request_concurrency = state
.distributed_request_concurrency_snapshot()
.await
.ok()
.flatten()
.map(|snapshot| {
serde_json::json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
Json(serde_json::json!({
"status": "ok",
"request_concurrency": request_concurrency,
"distributed_request_concurrency": distributed_request_concurrency,
}))
}
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
Json(state.hub.stats())
}
async fn metrics(State(state): State<AppState>) -> impl IntoResponse {
prometheus_response(&state.metric_samples().await)
}
pub async fn ws_proxy(
ws: WebSocketUpgrade,
State(state): State<AppState>,
headers: axum::http::HeaderMap,
) -> impl IntoResponse {
let node_id = headers
.get("x-node-id")
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.trim()
.to_string();
let node_name = headers
.get("x-node-name")
.and_then(|v| v.to_str().ok())
.unwrap_or(&node_id)
.trim()
.to_string();
let max_streams: usize = headers
.get("x-tunnel-max-streams")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse().ok())
.unwrap_or(state.max_streams)
.clamp(64, 2048);
if node_id.is_empty() {
warn!("proxy connection rejected: missing X-Node-ID header");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
Err(RequestAdmissionError::Local(ConcurrencyError::Saturated { .. }))
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Saturated {
..
}))
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Unavailable {
..
})) => return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response(),
Err(RequestAdmissionError::Local(ConcurrencyError::Closed { gate })) => {
warn!(
gate = gate,
"standalone tunnel relay request concurrency gate is closed"
);
return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response();
}
Err(RequestAdmissionError::Distributed(
DistributedConcurrencyError::InvalidConfiguration(message),
)) => {
warn!(
error = %message,
"standalone tunnel relay distributed request gate is invalid"
);
return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response();
}
};
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
hold_admission_permit_until(request_permit, async move {
proxy_conn::handle_proxy_connection(
socket,
state.hub,
node_id,
node_name,
max_streams,
state.proxy_conn_cfg,
)
.await
})
})
.into_response()
}
@@ -0,0 +1,14 @@
use bytes::Bytes;
pub use aether_contracts::tunnel::{
decode_payload, encode_frame, encode_goaway, encode_ping, encode_pong, encode_stream_error,
frame_payload_by_header, FrameHeader, RequestMeta, ResponseMeta, FLAG_END_STREAM,
FLAG_GZIP_COMPRESSED, GOAWAY, HEADER_SIZE, HEARTBEAT_ACK, HEARTBEAT_DATA, PING, PONG,
REQUEST_BODY, REQUEST_HEADERS, RESPONSE_BODY, RESPONSE_HEADERS, STREAM_END, STREAM_ERROR,
};
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
let (compressed, flags) =
aether_contracts::tunnel::compress_payload(Bytes::copy_from_slice(payload));
Ok((compressed.to_vec(), flags))
}
@@ -0,0 +1,157 @@
/// Proxy-side WebSocket connection handler
///
/// Handles the lifecycle of a single aether-proxy connection:
/// accept -> authenticate (headers) -> read loop -> cleanup
use std::sync::Arc;
use std::time::Duration;
use aether_runtime::bounded_queue;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::watch;
use tracing::{debug, info, warn};
use super::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
use super::protocol;
/// Maximum single frame size: 64 MB
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
pub async fn handle_proxy_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
node_id: String,
node_name: String,
max_streams: usize,
cfg: ConnConfig,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
let (close_tx, mut close_rx) = watch::channel(false);
let conn = Arc::new(ProxyConn::new(
conn_id,
node_id.clone(),
node_name.clone(),
tx,
close_tx,
max_streams,
));
hub.register_proxy(conn.clone());
let writer = tokio::spawn(async move {
loop {
tokio::select! {
msg = rx.recv() => match msg {
Some(msg) => {
if ws_tx.send(msg).await.is_err() {
break;
}
}
None => break,
},
changed = close_rx.changed() => {
if changed.is_err() || *close_rx.borrow() {
break;
}
}
}
}
let _ = ws_tx.close().await;
});
let ping_conn = conn.clone();
let ping_interval = cfg.ping_interval;
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if !matches!(
ping_conn.send(Message::Binary(ping.into())),
SendStatus::Queued
) {
break;
}
}
});
let reader_hub = hub.clone();
let reader_conn = conn.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
});
let _ = reader.await;
ping_task.abort();
conn.request_close();
hub.unregister_proxy(conn_id, &node_id);
drop(conn);
tokio::time::sleep(Duration::from_millis(100)).await;
writer.abort();
let _ = writer.await;
}
async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn: Arc<ProxyConn>,
idle_timeout: Duration,
) {
let idle_enabled = !idle_timeout.is_zero();
let mut oversized_count = 0u32;
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;
}
}
} else {
ws_rx.next().await
};
match msg {
Some(Ok(Message::Binary(data))) => {
let mut data = data.to_vec();
if data.len() > MAX_FRAME_SIZE {
oversized_count += 1;
warn!(
conn_id = conn.id,
size = data.len(),
"oversized frame from proxy"
);
if oversized_count >= 5 {
warn!(conn_id = conn.id, "too many oversized frames, closing");
conn.request_close();
break;
}
continue;
}
oversized_count = 0;
if data.len() < protocol::HEADER_SIZE {
debug!(conn_id = conn.id, "frame too small, skipping");
continue;
}
hub.handle_proxy_frame(conn.id, &mut data).await;
}
Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
break;
}
_ => {}
}
}
}