mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
Complete secure tunnel encryption support
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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!(
|
||||
|
||||
Reference in New Issue
Block a user