feat(tunnel): 引入 aether-hub 帧路由器,支持多 worker 共享 tunnel 连接

新增 Rust 实现的 aether-hub 服务,作为 Docker 容器内部 WebSocket 帧路由器,
解决多 Gunicorn worker 进程间 tunnel 连接隔离问题。

主要改动:
- 新增 aether-hub Rust 项目,实现 proxy/worker 双向帧路由与 stream_id 重映射
- 新增 HubConnectionManager/HubTunnelTransport,worker 通过 Hub 转发 tunnel 帧
- 新增 create_tunnel_transport 工厂函数,按运行环境自动选择 Hub 或直连模式
- 新增 NODE_STATUS 广播机制,Hub 实时通知所有 worker 节点连接状态变化
- CI/CD 新增 build-hub job,Dockerfile 集成 Hub 二进制,deploy.sh 适配 Hub 构建
- 默认 GUNICORN_WORKERS 从 4 降为 2
This commit is contained in:
fawney19
2026-03-02 02:43:14 +08:00
parent 97d42703da
commit 039a18c243
30 changed files with 3728 additions and 75 deletions

3
aether-hub/.dockerignore Normal file
View File

@@ -0,0 +1,3 @@
target/
.git/
.DS_Store

1229
aether-hub/Cargo.lock generated Normal file

File diff suppressed because it is too large Load Diff

23
aether-hub/Cargo.toml Normal file
View File

@@ -0,0 +1,23 @@
[package]
name = "aether-hub"
version = "0.1.0"
edition = "2021"
description = "Tunnel Hub for Aether - frame router between workers and proxies"
[dependencies]
tokio = { version = "1", features = ["full"] }
axum = { version = "0.8", features = ["ws"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
clap = { version = "4", features = ["derive", "env"] }
dashmap = "6"
parking_lot = "0.12"
flate2 = "1"
futures-util = "0.3"
[profile.release]
lto = true
strip = true
codegen-units = 1

27
aether-hub/Dockerfile Normal file
View File

@@ -0,0 +1,27 @@
# syntax=docker/dockerfile:1
FROM rust:1.85-slim AS builder
WORKDIR /build/aether-hub
# 先构建依赖层,最大化后续代码变更时的缓存命中
COPY Cargo.toml Cargo.lock ./
RUN mkdir src && printf 'fn main() {}\n' > src/main.rs
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked
RUN rm -rf src
COPY src ./src
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
cargo build --release --locked && \
cp target/release/aether-hub /tmp/aether-hub
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \
rm -rf /var/lib/apt/lists/*
COPY --from=builder /tmp/aether-hub /usr/local/bin/aether-hub
EXPOSE 8085
ENTRYPOINT ["/usr/local/bin/aether-hub"]
CMD ["--bind", "0.0.0.0:8085"]

695
aether-hub/src/hub.rs Normal file
View File

@@ -0,0 +1,695 @@
/// HubRouter -- central frame routing engine
///
/// Manages proxy connections (node_id -> [ProxyConn]) and worker connections (conn_id -> WorkerConn).
/// Routes frames between workers and proxies with stream_id remapping.
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use axum::extract::ws::Message;
use dashmap::DashMap;
use parking_lot::RwLock;
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::protocol;
// ---------------------------------------------------------------------------
// Proxy connection
// ---------------------------------------------------------------------------
pub struct ProxyConn {
pub id: u64,
pub node_id: String,
pub node_name: String,
pub tx: mpsc::UnboundedSender<Message>,
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: mpsc::UnboundedSender<Message>,
max_streams: usize,
) -> Self {
Self {
id,
node_id,
node_name,
tx,
next_stream_id: AtomicU32::new(2), // even IDs, start at 2
stream_count: AtomicUsize::new(0),
max_streams,
}
}
/// Allocate a proxy-side stream_id (even numbers)
pub fn alloc_stream_id(&self) -> Option<u32> {
// Reserve one stream slot first (CAS to honor max_streams under contention).
let mut current = self.stream_count.load(Ordering::Relaxed);
loop {
if current >= self.max_streams {
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 send(&self, msg: Message) -> bool {
self.tx.send(msg).is_ok()
}
}
// ---------------------------------------------------------------------------
// Worker connection
// ---------------------------------------------------------------------------
pub struct WorkerConn {
pub id: u64,
pub tx: mpsc::UnboundedSender<Message>,
}
impl WorkerConn {
pub fn new(id: u64, tx: mpsc::UnboundedSender<Message>) -> Self {
Self { id, tx }
}
pub fn send(&self, msg: Message) -> bool {
self.tx.send(msg).is_ok()
}
}
// ---------------------------------------------------------------------------
// Stream mapping entry
// ---------------------------------------------------------------------------
#[derive(Debug, Clone, Copy)]
struct ProxySide {
proxy_conn_id: u64,
proxy_stream_id: u32,
}
#[derive(Debug, Clone, Copy)]
struct WorkerSide {
worker_conn_id: u64,
worker_stream_id: u32,
}
// ---------------------------------------------------------------------------
// HubRouter
// ---------------------------------------------------------------------------
pub struct HubRouter {
/// node_id -> list of proxy connections
proxy_conns: RwLock<std::collections::HashMap<String, Vec<Arc<ProxyConn>>>>,
/// proxy_conn_id -> Arc<ProxyConn> (for reverse lookup)
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
/// worker_conn_id -> Arc<WorkerConn>
worker_conns: DashMap<u64, Arc<WorkerConn>>,
/// (worker_conn_id, worker_stream_id) -> ProxySide
worker_to_proxy: DashMap<(u64, u32), ProxySide>,
/// (proxy_conn_id, proxy_stream_id) -> WorkerSide
proxy_to_worker: DashMap<(u64, u32), WorkerSide>,
/// Connection ID generator
next_conn_id: AtomicU64,
/// Round-robin counter for heartbeat forwarding
heartbeat_rr: AtomicU64,
/// Heartbeat tag -> proxy_conn_id mapping (u32 tag fits in stream_id field)
heartbeat_tags: DashMap<u32, u64>,
/// Next heartbeat tag (wrapping u32)
next_heartbeat_tag: AtomicU32,
}
impl HubRouter {
pub fn new() -> Arc<Self> {
Arc::new(Self {
proxy_conns: RwLock::new(std::collections::HashMap::new()),
proxy_conns_by_id: DashMap::new(),
worker_conns: DashMap::new(),
worker_to_proxy: DashMap::new(),
proxy_to_worker: DashMap::new(),
next_conn_id: AtomicU64::new(1),
heartbeat_rr: AtomicU64::new(0),
heartbeat_tags: DashMap::new(),
next_heartbeat_tag: AtomicU32::new(1),
})
}
pub fn alloc_conn_id(&self) -> u64 {
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
}
// -----------------------------------------------------------------------
// Proxy connection management
// -----------------------------------------------------------------------
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 mut map = self.proxy_conns.write();
map.entry(node_id.clone()).or_default().push(conn);
let pool_size = 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"
);
drop(map);
self.broadcast_node_status(&node_id);
}
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
self.proxy_conns_by_id.remove(&conn_id);
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);
}
}
let pool_size = 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"
);
drop(map);
// Cancel all in-flight streams on this proxy connection
self.cancel_streams_for_proxy(conn_id);
self.broadcast_node_status(node_id);
}
/// Get least-loaded proxy connection for a node
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()
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
.cloned()
}
/// Get pool size for a node
fn proxy_conn_count(&self, node_id: &str) -> usize {
let map = self.proxy_conns.read();
map.get(node_id).map(|v| v.len()).unwrap_or(0)
}
// -----------------------------------------------------------------------
// Worker connection management
// -----------------------------------------------------------------------
pub fn register_worker(&self, conn: Arc<WorkerConn>) {
info!(worker_id = conn.id, "worker connected");
self.worker_conns.insert(conn.id, conn);
}
pub fn unregister_worker(&self, conn_id: u64) {
self.worker_conns.remove(&conn_id);
info!(worker_id = conn_id, "worker disconnected");
// Clean up all stream mappings for this worker
let to_remove: Vec<(u64, u32)> = self
.worker_to_proxy
.iter()
.filter(|e| e.key().0 == conn_id)
.map(|e| *e.key())
.collect();
for key in &to_remove {
if let Some((_, proxy_side)) = self.worker_to_proxy.remove(key) {
self.proxy_to_worker
.remove(&(proxy_side.proxy_conn_id, proxy_side.proxy_stream_id));
// Release stream count on proxy side
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_side.proxy_conn_id) {
pc.release_stream();
}
}
}
if !to_remove.is_empty() {
debug!(
worker_id = conn_id,
streams_cleaned = to_remove.len(),
"cleaned up worker streams"
);
}
}
// -----------------------------------------------------------------------
// Frame routing: Worker -> Proxy
// -----------------------------------------------------------------------
/// Handle a frame from a worker. Returns error message if routing fails.
pub fn handle_worker_frame(&self, worker_conn_id: u64, data: &mut [u8]) -> Option<String> {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return Some("invalid frame".to_string()),
};
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
if data.len() < expected_len {
return Some("incomplete frame payload".to_string());
}
match header.msg_type {
protocol::REQUEST_HEADERS => {
self.route_request_headers(worker_conn_id, header.stream_id, data)
}
protocol::REQUEST_BODY => {
if header.flags & protocol::FLAG_END_STREAM != 0 {
debug!(
worker_conn_id = worker_conn_id,
stream_id = header.stream_id,
"worker sent REQUEST_BODY with END_STREAM"
);
}
self.route_worker_to_proxy(worker_conn_id, header.stream_id, data, false);
None
}
protocol::STREAM_END | protocol::STREAM_ERROR => {
self.route_worker_to_proxy(worker_conn_id, header.stream_id, data, true);
None
}
protocol::GOAWAY => {
warn!(
worker_conn_id = worker_conn_id,
"received GOAWAY from worker connection"
);
None
}
protocol::PING => {
let payload = protocol::frame_payload(data).to_vec();
let pong = protocol::encode_pong(&payload);
if let Some(wc) = self.worker_conns.get(&worker_conn_id) {
let _ = wc.send(Message::Binary(pong.into()));
}
None
}
protocol::PONG => None, // Worker responded to our ping, nothing to do
_ => {
debug!(
msg_type = header.msg_type,
"unexpected frame type from worker"
);
None
}
}
}
/// Route REQUEST_HEADERS: extract node_id, allocate proxy stream, create mapping
fn route_request_headers(
&self,
worker_conn_id: u64,
worker_stream_id: u32,
data: &mut [u8],
) -> Option<String> {
// Parse payload to extract node_id, and pre-build frame with node_id stripped.
// stream_id is set to 0 first; we'll rewrite to proxy_stream_id after allocation.
let extracted = match protocol::rebuild_request_headers_without_node_id(data, 0) {
Ok(v) => v,
Err(e) => return Some(e),
};
let node_id = extracted.node_id;
// Find a proxy connection for this node
let proxy_conn = match self.get_proxy_conn(&node_id) {
Some(c) => c,
None => {
return Some(format!("no proxy connection for node {}", node_id));
}
};
// Allocate proxy-side stream_id
let proxy_stream_id = match proxy_conn.alloc_stream_id() {
Some(sid) => sid,
None => {
return Some(format!("stream limit reached for node {}", node_id));
}
};
let mut rebuilt_frame = extracted.rebuilt_frame;
protocol::rewrite_stream_id(&mut rebuilt_frame, proxy_stream_id);
// Record bidirectional mapping
self.worker_to_proxy.insert(
(worker_conn_id, worker_stream_id),
ProxySide {
proxy_conn_id: proxy_conn.id,
proxy_stream_id,
},
);
self.proxy_to_worker.insert(
(proxy_conn.id, proxy_stream_id),
WorkerSide {
worker_conn_id,
worker_stream_id,
},
);
if !proxy_conn.send(Message::Binary(rebuilt_frame.into())) {
// Send failed, clean up mapping
self.worker_to_proxy
.remove(&(worker_conn_id, worker_stream_id));
self.proxy_to_worker
.remove(&(proxy_conn.id, proxy_stream_id));
proxy_conn.release_stream();
return Some("proxy connection send failed".to_string());
}
None
}
/// Route non-header frames from worker to proxy (REQUEST_BODY etc.)
fn route_worker_to_proxy(
&self,
worker_conn_id: u64,
worker_stream_id: u32,
data: &mut [u8],
terminal: bool,
) {
let proxy_side = if terminal {
match self
.worker_to_proxy
.remove(&(worker_conn_id, worker_stream_id))
{
Some((_, ps)) => {
self.proxy_to_worker
.remove(&(ps.proxy_conn_id, ps.proxy_stream_id));
if let Some(pc) = self.proxy_conns_by_id.get(&ps.proxy_conn_id) {
pc.release_stream();
}
ps
}
None => return, // Silently discard -- mapping already removed (race condition)
}
} else {
match self
.worker_to_proxy
.get(&(worker_conn_id, worker_stream_id))
{
Some(entry) => *entry.value(),
None => return, // Silently discard -- mapping already removed (race condition)
}
};
// Rewrite stream_id
protocol::rewrite_stream_id(data, proxy_side.proxy_stream_id);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_side.proxy_conn_id) {
let _ = pc.send(Message::Binary(data.to_vec().into()));
}
}
// -----------------------------------------------------------------------
// Frame routing: Proxy -> Worker
// -----------------------------------------------------------------------
/// Handle a frame from a proxy connection
pub 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 | protocol::RESPONSE_BODY => {
self.route_proxy_to_worker(proxy_conn_id, header.stream_id, data, false);
}
_ if header.is_stream_terminal() => {
self.route_proxy_to_worker(proxy_conn_id, header.stream_id, data, true);
}
protocol::HEARTBEAT_DATA => {
self.forward_heartbeat_to_worker(proxy_conn_id, data);
}
protocol::PONG => {} // Proxy responded to our ping
protocol::GOAWAY => {
warn!(
proxy_conn_id = proxy_conn_id,
"received GOAWAY from proxy connection"
);
}
protocol::PING => {
// Proxy sent a ping, reply with pong
let payload = if data.len() > protocol::HEADER_SIZE {
&data[protocol::HEADER_SIZE..]
} else {
&[]
};
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()));
}
}
_ => {
debug!(
msg_type = header.msg_type,
proxy_conn_id = proxy_conn_id,
"unexpected frame type from proxy"
);
}
}
}
/// Route response frames from proxy to worker
fn route_proxy_to_worker(
&self,
proxy_conn_id: u64,
proxy_stream_id: u32,
data: &mut [u8],
terminal: bool,
) {
let worker_side = if terminal {
// Remove mapping on terminal frames
match self
.proxy_to_worker
.remove(&(proxy_conn_id, proxy_stream_id))
{
Some((_, ws)) => {
self.worker_to_proxy
.remove(&(ws.worker_conn_id, ws.worker_stream_id));
// Release stream count
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
pc.release_stream();
}
ws
}
None => return, // Silently discard
}
} else {
match self.proxy_to_worker.get(&(proxy_conn_id, proxy_stream_id)) {
Some(entry) => *entry.value(),
None => return, // Silently discard
}
};
// Rewrite stream_id to worker-side
protocol::rewrite_stream_id(data, worker_side.worker_stream_id);
if let Some(wc) = self.worker_conns.get(&worker_side.worker_conn_id) {
let _ = wc.send(Message::Binary(data.to_vec().into()));
}
}
/// Forward HEARTBEAT_DATA to a worker (round-robin)
fn forward_heartbeat_to_worker(&self, proxy_conn_id: u64, data: &[u8]) {
// Pick a worker via round-robin
let workers: Vec<Arc<WorkerConn>> = self
.worker_conns
.iter()
.map(|e| e.value().clone())
.collect();
if workers.is_empty() {
debug!("no workers to forward heartbeat to");
return;
}
let idx = self.heartbeat_rr.fetch_add(1, Ordering::Relaxed) as usize % workers.len();
let worker = &workers[idx];
// Use a u32 tag in the stream_id field to identify the proxy connection.
// The tag maps to the full u64 proxy_conn_id via heartbeat_tags DashMap,
// avoiding truncation of u64 conn_id to u32.
// Skip 0 (reserved for control frames) via CAS loop.
let tag = loop {
let t = self.next_heartbeat_tag.fetch_add(1, Ordering::Relaxed);
if t != 0 {
break t;
}
};
self.heartbeat_tags.insert(tag, proxy_conn_id);
let mut forwarded = data.to_vec();
protocol::rewrite_stream_id(&mut forwarded, tag);
let _ = worker.send(Message::Binary(forwarded.into()));
}
/// Handle HEARTBEAT_ACK from worker -- route back to the proxy
pub fn handle_worker_heartbeat_ack(&self, data: &mut [u8]) {
let header = match protocol::FrameHeader::parse(data) {
Some(h) => h,
None => return,
};
// Recover the original proxy_conn_id from the tag stored in stream_id
let tag = header.stream_id;
let proxy_conn_id = match self.heartbeat_tags.remove(&tag) {
Some((_, id)) => id,
None => return,
};
// Reset stream_id to 0 before forwarding to proxy
protocol::rewrite_stream_id(data, 0);
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
let _ = pc.send(Message::Binary(data.to_vec().into()));
}
}
// -----------------------------------------------------------------------
// Stream cleanup
// -----------------------------------------------------------------------
/// Cancel all in-flight streams for a disconnected proxy connection
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
let to_remove: Vec<((u64, u32), WorkerSide)> = self
.proxy_to_worker
.iter()
.filter(|e| e.key().0 == proxy_conn_id)
.map(|e| (*e.key(), *e.value()))
.collect();
for ((p_conn_id, p_sid), worker_side) in &to_remove {
self.proxy_to_worker.remove(&(*p_conn_id, *p_sid));
self.worker_to_proxy
.remove(&(worker_side.worker_conn_id, worker_side.worker_stream_id));
// Send STREAM_ERROR to worker
let err_frame =
protocol::encode_stream_error(worker_side.worker_stream_id, "proxy disconnected");
if let Some(wc) = self.worker_conns.get(&worker_side.worker_conn_id) {
let _ = wc.send(Message::Binary(err_frame.into()));
}
}
if !to_remove.is_empty() {
warn!(
proxy_conn_id = proxy_conn_id,
streams_cancelled = to_remove.len(),
"cancelled in-flight streams due to proxy disconnect"
);
}
}
// -----------------------------------------------------------------------
// NODE_STATUS broadcast
// -----------------------------------------------------------------------
fn broadcast_node_status(&self, node_id: &str) {
let conn_count = self.proxy_conn_count(node_id);
let connected = conn_count > 0;
let frame = protocol::encode_node_status(node_id, connected, conn_count);
let msg = Message::Binary(frame.into());
let mut sent = 0usize;
for entry in self.worker_conns.iter() {
if entry.value().send(msg.clone()) {
sent += 1;
}
}
debug!(
node_id = %node_id,
connected = connected,
conn_count = conn_count,
workers_notified = sent,
"broadcast NODE_STATUS"
);
}
// -----------------------------------------------------------------------
// Stats
// -----------------------------------------------------------------------
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,
worker_connections: self.worker_conns.len(),
nodes,
active_streams: self.worker_to_proxy.len(),
}
}
}
#[derive(serde::Serialize)]
pub struct HubStats {
pub proxy_connections: usize,
pub worker_connections: usize,
pub nodes: usize,
pub active_streams: usize,
}

159
aether-hub/src/main.rs Normal file
View File

@@ -0,0 +1,159 @@
mod hub;
mod protocol;
mod proxy_conn;
mod worker_conn;
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State;
use axum::response::{IntoResponse, Json};
use axum::routing::get;
use axum::Router;
use clap::Parser;
use tracing::{info, warn};
use crate::hub::HubRouter;
#[derive(Parser, Debug)]
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
struct Args {
/// Bind address
#[arg(long, default_value = "0.0.0.0:8085", env = "TUNNEL_HUB_BIND")]
bind: String,
/// Proxy-side idle timeout in seconds
#[arg(long, default_value_t = 90, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
proxy_idle_timeout: u64,
/// Worker-side idle timeout in seconds
#[arg(long, default_value_t = 60, env = "TUNNEL_HUB_WORKER_IDLE_TIMEOUT")]
worker_idle_timeout: u64,
/// Ping interval in seconds (for both sides)
#[arg(long, default_value_t = 15, env = "TUNNEL_HUB_PING_INTERVAL")]
ping_interval: u64,
/// Max concurrent streams per proxy connection
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
max_streams: usize,
}
#[derive(Clone)]
struct AppState {
hub: Arc<HubRouter>,
proxy_idle_timeout: Duration,
worker_idle_timeout: Duration,
ping_interval: Duration,
max_streams: usize,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// Initialize tracing
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "aether_hub=info".into()),
)
.init();
let args = Args::parse();
let hub = HubRouter::new();
let state = AppState {
hub,
proxy_idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
worker_idle_timeout: Duration::from_secs(args.worker_idle_timeout),
ping_interval: Duration::from_secs(args.ping_interval),
max_streams: args.max_streams,
};
let app = Router::new()
.route("/health", get(health))
.route("/stats", get(stats))
.route("/proxy", get(ws_proxy))
.route("/worker", get(ws_worker))
.with_state(state);
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
info!(bind = %args.bind, "aether-hub started");
axum::serve(listener, app).await?;
Ok(())
}
// ---------------------------------------------------------------------------
// HTTP endpoints
// ---------------------------------------------------------------------------
async fn health() -> impl IntoResponse {
Json(serde_json::json!({"status": "ok"}))
}
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
Json(state.hub.stats())
}
// ---------------------------------------------------------------------------
// WebSocket endpoints
// ---------------------------------------------------------------------------
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();
}
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
proxy_conn::handle_proxy_connection(
socket,
state.hub,
node_id,
node_name,
max_streams,
state.ping_interval,
state.proxy_idle_timeout,
)
})
.into_response()
}
async fn ws_worker(ws: WebSocketUpgrade, State(state): State<AppState>) -> impl IntoResponse {
ws.max_frame_size(64 * 1024 * 1024)
.on_upgrade(move |socket| {
worker_conn::handle_worker_connection(
socket,
state.hub,
state.ping_interval,
state.worker_idle_timeout,
)
})
}

232
aether-hub/src/protocol.rs Normal file
View File

@@ -0,0 +1,232 @@
/// Tunnel binary frame protocol
///
/// Frame format (10-byte header + payload):
/// | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
use std::io::Read;
use flate2::read::GzDecoder;
use flate2::write::GzEncoder;
use flate2::Compression;
pub const HEADER_SIZE: usize = 10;
// Message types
pub const REQUEST_HEADERS: u8 = 0x01;
pub const REQUEST_BODY: u8 = 0x02;
pub const RESPONSE_HEADERS: u8 = 0x03;
pub const RESPONSE_BODY: u8 = 0x04;
pub const STREAM_END: u8 = 0x05;
pub const STREAM_ERROR: u8 = 0x06;
pub const PING: u8 = 0x10;
pub const PONG: u8 = 0x11;
pub const GOAWAY: u8 = 0x12;
pub const HEARTBEAT_DATA: u8 = 0x13;
pub const HEARTBEAT_ACK: u8 = 0x14;
pub const NODE_STATUS: u8 = 0x15;
// Flags
pub const FLAG_END_STREAM: u8 = 0x01;
pub const FLAG_GZIP_COMPRESSED: u8 = 0x02;
#[derive(Debug, Clone, Copy)]
pub struct FrameHeader {
pub stream_id: u32,
pub msg_type: u8,
pub flags: u8,
pub payload_len: u32,
}
impl FrameHeader {
/// Parse frame header from raw bytes (must be >= HEADER_SIZE)
#[inline]
pub fn parse(data: &[u8]) -> Option<Self> {
if data.len() < HEADER_SIZE {
return None;
}
Some(Self {
stream_id: u32::from_be_bytes([data[0], data[1], data[2], data[3]]),
msg_type: data[4],
flags: data[5],
payload_len: u32::from_be_bytes([data[6], data[7], data[8], data[9]]),
})
}
/// Check if this is a stream-terminating frame
#[inline]
pub fn is_stream_terminal(&self) -> bool {
self.msg_type == STREAM_END || self.msg_type == STREAM_ERROR
}
}
#[derive(Debug)]
pub struct RequestHeadersExtracted {
pub node_id: String,
pub rebuilt_frame: Vec<u8>,
}
/// Encode a STREAM_ERROR frame for a given stream_id with an error message
pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
let payload = msg.as_bytes();
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
buf.extend_from_slice(&stream_id.to_be_bytes());
buf.push(STREAM_ERROR);
buf.push(0); // flags
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
buf.extend_from_slice(payload);
buf
}
/// Encode a NODE_STATUS frame (stream_id=0, Hub-generated)
pub fn encode_node_status(node_id: &str, connected: bool, conn_count: usize) -> Vec<u8> {
let payload = serde_json::json!({
"node_id": node_id,
"connected": connected,
"conn_count": conn_count,
});
let payload_bytes = payload.to_string().into_bytes();
let mut buf = Vec::with_capacity(HEADER_SIZE + payload_bytes.len());
buf.extend_from_slice(&0u32.to_be_bytes()); // stream_id = 0
buf.push(NODE_STATUS);
buf.push(0); // flags
buf.extend_from_slice(&(payload_bytes.len() as u32).to_be_bytes());
buf.extend_from_slice(&payload_bytes);
buf
}
/// Encode a PING frame (stream_id=0)
pub fn encode_ping() -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE);
buf.extend_from_slice(&0u32.to_be_bytes());
buf.push(PING);
buf.push(0);
buf.extend_from_slice(&0u32.to_be_bytes());
buf
}
/// Encode a PONG frame (stream_id=0, echo payload)
pub fn encode_pong(payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
buf.extend_from_slice(&0u32.to_be_bytes());
buf.push(PONG);
buf.push(0);
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
buf.extend_from_slice(payload);
buf
}
/// Encode a GOAWAY frame (stream_id=0)
pub fn encode_goaway() -> Vec<u8> {
let mut buf = Vec::with_capacity(HEADER_SIZE);
buf.extend_from_slice(&0u32.to_be_bytes());
buf.push(GOAWAY);
buf.push(0);
buf.extend_from_slice(&0u32.to_be_bytes());
buf
}
/// Rewrite the stream_id in raw frame bytes (first 4 bytes) -- near zero-copy
#[inline]
pub fn rewrite_stream_id(data: &mut [u8], new_stream_id: u32) {
let bytes = new_stream_id.to_be_bytes();
data[0] = bytes[0];
data[1] = bytes[1];
data[2] = bytes[2];
data[3] = bytes[3];
}
/// Get the payload portion of a raw frame (after the 10-byte header)
#[inline]
pub fn frame_payload(data: &[u8]) -> &[u8] {
if data.len() > HEADER_SIZE {
&data[HEADER_SIZE..]
} else {
&[]
}
}
/// Parse REQUEST_HEADERS payload, extract `node_id`, strip it from JSON,
/// and rebuild a new REQUEST_HEADERS frame with `new_stream_id`.
///
/// If the source frame is gzip-compressed, this function will decode it first,
/// then try to re-encode with gzip (only keeps compression when payload shrinks).
pub fn rebuild_request_headers_without_node_id(
data: &[u8],
new_stream_id: u32,
) -> Result<RequestHeadersExtracted, String> {
let header = FrameHeader::parse(data).ok_or_else(|| "invalid frame header".to_string())?;
if header.msg_type != REQUEST_HEADERS {
return Err("frame is not REQUEST_HEADERS".to_string());
}
let payload = frame_payload_by_header(data, &header)
.ok_or_else(|| "incomplete REQUEST_HEADERS payload".to_string())?;
let decoded_payload = if header.flags & FLAG_GZIP_COMPRESSED != 0 {
let mut decoder = GzDecoder::new(payload);
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.map_err(|e| format!("failed to decompress REQUEST_HEADERS: {e}"))?;
decoded
} else {
payload.to_vec()
};
let mut meta: serde_json::Value = serde_json::from_slice(&decoded_payload)
.map_err(|e| format!("invalid REQUEST_HEADERS JSON: {e}"))?;
let obj = meta
.as_object_mut()
.ok_or_else(|| "REQUEST_HEADERS payload must be a JSON object".to_string())?;
let node_id = obj
.remove("node_id")
.and_then(|v| v.as_str().map(|s| s.to_string()))
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.ok_or_else(|| "missing node_id in REQUEST_HEADERS".to_string())?;
let stripped_payload = serde_json::to_vec(&meta)
.map_err(|e| format!("failed to encode REQUEST_HEADERS payload: {e}"))?;
let (final_payload, flags) =
maybe_recompress_payload(&stripped_payload, header.flags & FLAG_GZIP_COMPRESSED != 0)
.map_err(|e| format!("failed to recompress REQUEST_HEADERS payload: {e}"))?;
let mut rebuilt = Vec::with_capacity(HEADER_SIZE + final_payload.len());
rebuilt.extend_from_slice(&new_stream_id.to_be_bytes());
rebuilt.push(REQUEST_HEADERS);
rebuilt.push(flags);
rebuilt.extend_from_slice(&(final_payload.len() as u32).to_be_bytes());
rebuilt.extend_from_slice(&final_payload);
Ok(RequestHeadersExtracted {
node_id,
rebuilt_frame: rebuilt,
})
}
#[inline]
fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
let payload_len = header.payload_len as usize;
let end = HEADER_SIZE.checked_add(payload_len)?;
if data.len() < end {
return None;
}
Some(&data[HEADER_SIZE..end])
}
fn maybe_recompress_payload(
payload: &[u8],
prefer_gzip: bool,
) -> Result<(Vec<u8>, u8), std::io::Error> {
if !prefer_gzip {
return Ok((payload.to_vec(), 0));
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
std::io::Write::write_all(&mut encoder, payload)?;
let compressed = encoder.finish()?;
if compressed.len() < payload.len() {
Ok((compressed, FLAG_GZIP_COMPRESSED))
} else {
Ok((payload.to_vec(), 0))
}
}

View File

@@ -0,0 +1,144 @@
/// 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 axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::hub::{HubRouter, ProxyConn};
use crate::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,
ping_interval: Duration,
idle_timeout: Duration,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let conn = Arc::new(ProxyConn::new(
conn_id,
node_id.clone(),
node_name.clone(),
tx,
max_streams,
));
hub.register_proxy(conn.clone());
// Spawn writer task: drains channel -> WebSocket
let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
if ws_tx.send(msg).await.is_err() {
break;
}
}
let _ = ws_tx.close().await;
});
// Spawn ping task
let ping_tx = conn.tx.clone();
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if ping_tx.send(Message::Binary(ping.into())).is_err() {
break;
}
}
});
// Spawn reader task
let reader_hub = hub.clone();
let reader_node_id = node_id.clone();
let reader_tx = conn.tx.clone();
let reader = tokio::spawn(async move {
run_proxy_reader(
ws_rx,
reader_hub,
conn_id,
reader_node_id,
reader_tx,
idle_timeout,
)
.await;
});
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
let _ = reader.await;
ping_task.abort();
writer.abort();
hub.unregister_proxy(conn_id, &node_id);
}
async fn run_proxy_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn_id: u64,
node_id: String,
tx: mpsc::UnboundedSender<Message>,
idle_timeout: Duration,
) {
let mut oversized_count = 0u32;
loop {
let msg = tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(conn_id = conn_id, node_id = %node_id, "proxy idle timeout");
let _ = tx.send(Message::Binary(protocol::encode_goaway().into()));
break;
}
};
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");
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);
}
Some(Ok(Message::Close(_))) | None => {
info!(conn_id = conn_id, node_id = %node_id, "proxy WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(conn_id = conn_id, error = %e, "proxy WebSocket error");
break;
}
_ => {} // Ignore text/ping/pong at WS level
}
}
}

View File

@@ -0,0 +1,130 @@
/// Worker-side WebSocket connection handler
///
/// Handles the lifecycle of a single Gunicorn worker connection:
/// accept -> read loop (route frames via Hub) -> cleanup
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::hub::{HubRouter, WorkerConn};
use crate::protocol;
pub async fn handle_worker_connection(
ws: WebSocket,
hub: Arc<HubRouter>,
ping_interval: Duration,
idle_timeout: Duration,
) {
let conn_id = hub.alloc_conn_id();
let (mut ws_tx, ws_rx) = ws.split();
// Create channel for outbound messages
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let conn = Arc::new(WorkerConn::new(conn_id, tx));
hub.register_worker(conn.clone());
// Spawn writer task
let writer = tokio::spawn(async move {
while let Some(msg) = rx.recv().await {
if ws_tx.send(msg).await.is_err() {
break;
}
}
let _ = ws_tx.close().await;
});
// Spawn ping task
let ping_tx = conn.tx.clone();
let ping_task = tokio::spawn(async move {
loop {
tokio::time::sleep(ping_interval).await;
let ping = protocol::encode_ping();
if ping_tx.send(Message::Binary(ping.into())).is_err() {
break;
}
}
});
// Spawn reader task
let reader_hub = hub.clone();
let reader_tx = conn.tx.clone();
let reader = tokio::spawn(async move {
run_worker_reader(
ws_rx,
reader_hub,
conn_id,
conn.clone(),
reader_tx,
idle_timeout,
)
.await;
});
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
let _ = reader.await;
ping_task.abort();
writer.abort();
hub.unregister_worker(conn_id);
}
async fn run_worker_reader(
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
hub: Arc<HubRouter>,
conn_id: u64,
conn: Arc<WorkerConn>,
tx: mpsc::UnboundedSender<Message>,
idle_timeout: Duration,
) {
loop {
let msg = tokio::select! {
msg = ws_rx.next() => msg,
_ = tokio::time::sleep(idle_timeout) => {
warn!(worker_id = conn_id, "worker idle timeout");
let _ = tx.send(Message::Binary(protocol::encode_goaway().into()));
break;
}
};
match msg {
Some(Ok(Message::Binary(data))) => {
let mut data = data.to_vec();
if data.len() < protocol::HEADER_SIZE {
debug!(worker_id = conn_id, "frame too small, skipping");
continue;
}
let header = match protocol::FrameHeader::parse(&data) {
Some(h) => h,
None => continue,
};
// HEARTBEAT_ACK from worker -> route back to proxy
if header.msg_type == protocol::HEARTBEAT_ACK {
hub.handle_worker_heartbeat_ack(&mut data);
continue;
}
// Regular frames: route via hub
if let Some(err_msg) = hub.handle_worker_frame(conn_id, &mut data) {
// Send STREAM_ERROR back to worker
let err_frame = protocol::encode_stream_error(header.stream_id, &err_msg);
let _ = conn.send(Message::Binary(err_frame.into()));
}
}
Some(Ok(Message::Close(_))) | None => {
info!(worker_id = conn_id, "worker WebSocket closed");
break;
}
Some(Err(e)) => {
warn!(worker_id = conn_id, error = %e, "worker WebSocket error");
break;
}
_ => {} // Ignore text/ping/pong at WS level
}
}
}