mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
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:
@@ -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;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user