Complete secure tunnel encryption support

This commit is contained in:
RWDai
2026-05-22 09:47:28 +08:00
parent 4f49dd5943
commit 2e701a90c9
24 changed files with 720 additions and 33 deletions
+37 -6
View File
@@ -17,6 +17,10 @@ use crate::egress_proxy::{
};
use crate::state::{AppState, ServerContext};
use aether_contracts::tunnel::{CURRENT_TUNNEL_PROTOCOL_VERSION, TUNNEL_PROTOCOL_VERSION_HEADER};
use aether_contracts::tunnel_security::{
SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED,
TUNNEL_SECURITY_SESSION_HEADER,
};
use super::{dispatcher, heartbeat, writer};
@@ -45,16 +49,28 @@ pub async fn connect_and_run(
// Build WebSocket request with auth headers
let mut request = ws_url.clone().into_client_request()?;
let headers = request.headers_mut();
headers.insert(
"Authorization",
http::HeaderValue::from_str(&format!("Bearer {}", server.management_token))?,
);
if server.tunnel_security != crate::config::TunnelSecurity::NonTlsRequired {
headers.insert(
"Authorization",
http::HeaderValue::from_str(&format!("Bearer {}", server.management_token))?,
);
}
headers.insert(
TUNNEL_PROTOCOL_VERSION_HEADER,
http::HeaderValue::from_str(&CURRENT_TUNNEL_PROTOCOL_VERSION.to_string())?,
);
let node_id = server.node_id.read().unwrap().clone();
headers.insert("X-Node-Id", http::HeaderValue::from_str(&node_id)?);
if server.tunnel_security == crate::config::TunnelSecurity::NonTlsRequired {
headers.insert(
TUNNEL_SECURITY_HEADER,
http::HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED),
);
headers.insert(
TUNNEL_SECURITY_SESSION_HEADER,
http::HeaderValue::from_str(&node_id)?,
);
}
// Use dynamic node_name (may be updated by remote config) instead of
// the static server.node_name, so that remote name changes take effect
// on the next reconnect.
@@ -119,6 +135,19 @@ pub async fn connect_and_run(
handshake_timeout.as_millis()
)
})??;
let security = if server.tunnel_security == crate::config::TunnelSecurity::NonTlsRequired {
let key = server
.tunnel_encryption_key
.as_deref()
.ok_or_else(|| anyhow::anyhow!("secure tunnel requires tunnel_encryption_key"))?;
Some(Arc::new(SecureFrameCodec::new(
key,
&node_id,
TunnelSecurityRole::Client,
)?))
} else {
None
};
let stale_timeout = state
.config
.tunnel_stale_timeout()
@@ -146,10 +175,11 @@ pub async fn connect_and_run(
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
// Spawn writer task (with WebSocket ping keepalive)
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics(
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security(
ws_sink,
ping_interval,
Some(Arc::clone(&server.tunnel_metrics)),
security.clone(),
);
let drain_signal = spawn_drain_signal(conn_idx, frame_tx.clone(), drain.clone());
@@ -174,13 +204,14 @@ pub async fn connect_and_run(
let state_clone = Arc::clone(state);
let server_clone = Arc::clone(server);
let outcome = tokio::select! {
result = dispatcher::run(
result = dispatcher::run_with_security(
state_clone,
server_clone,
ws_read,
frame_tx.clone(),
hb_handle,
drain.clone(),
security.clone(),
) => {
match result {
Ok(()) => Ok(TunnelOutcome::Disconnected),
@@ -18,6 +18,7 @@ use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::stream_handler;
use super::writer::FrameSender;
use aether_contracts::tunnel_security::SecureFrameCodec;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StreamDispatchStatus {
@@ -27,13 +28,32 @@ enum StreamDispatchStatus {
}
/// Run the dispatcher loop, reading from the WebSocket stream.
#[allow(dead_code)]
pub async fn run<S>(
state: Arc<AppState>,
server: Arc<ServerContext>,
ws_stream: S,
frame_tx: FrameSender,
heartbeat: HeartbeatHandle,
drain: watch::Receiver<bool>,
) -> Result<(), anyhow::Error>
where
S: StreamExt<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
+ Unpin
+ Send
+ 'static,
{
run_with_security(state, server, ws_stream, frame_tx, heartbeat, drain, None).await
}
pub async fn run_with_security<S>(
state: Arc<AppState>,
server: Arc<ServerContext>,
mut ws_stream: S,
frame_tx: FrameSender,
heartbeat: HeartbeatHandle,
mut drain: watch::Receiver<bool>,
security: Option<Arc<SecureFrameCodec>>,
) -> Result<(), anyhow::Error>
where
S: StreamExt<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
@@ -130,6 +150,19 @@ where
continue;
}
};
let frame = match security.as_deref() {
Some(codec) => match codec.decrypt_frame(frame) {
Ok(frame) => frame,
Err(e) => {
warn!(error = %e, "failed to decrypt secure tunnel frame");
server
.tunnel_metrics
.record_error("secure_frame_decrypt_error", &e.to_string());
break None;
}
},
None => frame,
};
match frame.msg_type {
MsgType::RequestHeaders => {
@@ -425,6 +425,8 @@ mod tests {
server_label: "heartbeat-test".to_string(),
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
tunnel_security: config.tunnel_security,
tunnel_encryption_key: config.tunnel_encryption_key.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(RwLock::new("node-123".to_string())),
aether_client: Arc::new(AetherClient::new(
+2
View File
@@ -476,6 +476,8 @@ mod tests {
server_label: "gateway-owned-tunnel".to_string(),
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
tunnel_security: config.tunnel_security,
tunnel_encryption_key: config.tunnel_encryption_key.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(std::sync::RwLock::new(node_id.to_string())),
aether_client: Arc::new(AetherClient::new(
@@ -2491,6 +2491,8 @@ mod tests {
server_label: "server".to_string(),
aether_url: config.aether_url.clone(),
management_token: config.management_token.clone(),
tunnel_security: config.tunnel_security,
tunnel_encryption_key: config.tunnel_encryption_key.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(std::sync::RwLock::new("node-1".to_string())),
aether_client: Arc::new(AetherClient::new(
+43 -5
View File
@@ -20,6 +20,7 @@ use tracing::{debug, error, trace};
use crate::state::TunnelMetrics;
use super::protocol::Frame;
use aether_contracts::tunnel_security::SecureFrameCodec;
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
@@ -89,10 +90,23 @@ where
}
/// Spawn the writer task with optional tunnel metrics instrumentation.
#[allow(dead_code)]
pub fn spawn_writer_with_metrics<S>(
sink: S,
ping_interval: Duration,
tunnel_metrics: Option<Arc<TunnelMetrics>>,
) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
spawn_writer_with_metrics_and_security(sink, ping_interval, tunnel_metrics, None)
}
pub fn spawn_writer_with_metrics_and_security<S>(
mut sink: S,
ping_interval: Duration,
tunnel_metrics: Option<Arc<TunnelMetrics>>,
security: Option<Arc<SecureFrameCodec>>,
) -> (FrameSender, JoinHandle<()>)
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
@@ -109,7 +123,14 @@ where
loop {
if let Ok(frame) = high_rx.try_recv() {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
if !write_frame(
&mut sink,
frame,
tunnel_metrics.as_deref(),
security.as_deref(),
)
.await
{
break;
}
continue;
@@ -123,7 +144,7 @@ where
frame = high_rx.recv(), if high_open => {
match frame {
Some(frame) => {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref(), security.as_deref()).await {
break;
}
}
@@ -143,7 +164,7 @@ where
frame = normal_rx.recv(), if normal_open => {
match frame {
Some(frame) => {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref()).await {
if !write_frame(&mut sink, frame, tunnel_metrics.as_deref(), security.as_deref()).await {
break;
}
}
@@ -175,14 +196,31 @@ fn classify_frame_priority(frame: &Frame) -> FramePriority {
}
}
async fn write_frame<S>(sink: &mut S, frame: Frame, tunnel_metrics: Option<&TunnelMetrics>) -> bool
async fn write_frame<S>(
sink: &mut S,
frame: Frame,
tunnel_metrics: Option<&TunnelMetrics>,
security: Option<&SecureFrameCodec>,
) -> bool
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
{
let stream_id = frame.stream_id;
let msg_type = frame.msg_type;
let flags = frame.flags;
let data = frame.encode();
let data = match security {
Some(codec) => match codec.encrypt_frame(frame) {
Ok(data) => data,
Err(e) => {
error!(error = %e, "failed to encrypt tunnel frame");
if let Some(metrics) = tunnel_metrics {
metrics.record_error("secure_frame_encrypt_error", &e.to_string());
}
return false;
}
},
None => frame.encode(),
};
let wire_len = data.len().max(HEADER_SIZE);
if let Err(e) = sink.send(Message::Binary(data.into())).await {
error!(