mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat(security): harden gateway boundaries and usage policies
Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change. Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
@@ -1,12 +1,29 @@
|
||||
use aether_http::{build_http_client, HttpClientConfig};
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
use futures_util::future::BoxFuture;
|
||||
use reqwest::Client;
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_contracts::tunnel_security::{
|
||||
sign_tunnel_control_plane_request_for_generation, TUNNEL_CONTROL_PLANE_GENERATION_HEADER,
|
||||
TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, TUNNEL_CONTROL_PLANE_NONCE_HEADER,
|
||||
TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER,
|
||||
};
|
||||
use aether_gateway_tunnel::{TUNNEL_HEARTBEAT_PATH, TUNNEL_NODE_STATUS_PATH};
|
||||
|
||||
use super::hub::ProxyConn;
|
||||
|
||||
const MAX_CONTROL_PLANE_RESPONSE_BYTES: usize = 256 * 1024;
|
||||
const MAX_CONTROL_PLANE_BASE_URL_BYTES: usize = 2 * 1024;
|
||||
pub(crate) const CONTROL_PLANE_CREDENTIAL_REVOKED: &str = "proxy tunnel credential revoked";
|
||||
pub(crate) const CONTROL_PLANE_CREDENTIAL_UNAVAILABLE: &str =
|
||||
"proxy tunnel credential validation unavailable";
|
||||
|
||||
type HeartbeatAckCallback =
|
||||
dyn Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync;
|
||||
type NodeStatusCallback =
|
||||
dyn Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>> + Send + Sync;
|
||||
dyn Fn(Arc<ProxyConn>, Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync;
|
||||
type NodeStatusCallback = dyn Fn(Arc<ProxyConn>, bool, usize, u64) -> BoxFuture<'static, Result<(), String>>
|
||||
+ Send
|
||||
+ Sync;
|
||||
|
||||
enum ControlPlaneMode {
|
||||
Disabled,
|
||||
@@ -27,11 +44,32 @@ pub struct ControlPlaneClient {
|
||||
|
||||
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()
|
||||
})
|
||||
// The control-plane base URL is operator supplied, but it is used for
|
||||
// requests carrying a tunnel credential. Reject URL-controlled
|
||||
// request components (userinfo/query/fragment) before concatenating
|
||||
// endpoint paths; otherwise a typo such as `?token=...` can leak
|
||||
// credentials or change the signed request target. Keep the
|
||||
// established support for HTTP and private deployment hosts—the
|
||||
// standalone tunnel commonly talks to an in-cluster gateway.
|
||||
let Some(base_url) = normalize_control_plane_base_url(&base_url) else {
|
||||
return Self {
|
||||
inner: Arc::new(ControlPlaneMode::Http {
|
||||
client: None,
|
||||
base_url: String::new(),
|
||||
}),
|
||||
};
|
||||
};
|
||||
let client = apply_http_client_config(
|
||||
reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.redirect(reqwest::redirect::Policy::none()),
|
||||
&HttpClientConfig {
|
||||
request_timeout_ms: Some(10_000),
|
||||
user_agent: Some("aether-tunnel-standalone/control-plane".to_string()),
|
||||
..HttpClientConfig::default()
|
||||
},
|
||||
)
|
||||
.build()
|
||||
.ok();
|
||||
Self {
|
||||
inner: Arc::new(ControlPlaneMode::Http { client, base_url }),
|
||||
@@ -49,9 +87,11 @@ impl ControlPlaneClient {
|
||||
push_node_status: PushNodeStatus,
|
||||
) -> Self
|
||||
where
|
||||
HeartbeatAck:
|
||||
Fn(Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>> + Send + Sync + 'static,
|
||||
PushNodeStatus: Fn(String, bool, usize, u64) -> BoxFuture<'static, Result<(), String>>
|
||||
HeartbeatAck: Fn(Arc<ProxyConn>, Vec<u8>) -> BoxFuture<'static, Result<Vec<u8>, String>>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
PushNodeStatus: Fn(Arc<ProxyConn>, bool, usize, u64) -> BoxFuture<'static, Result<(), String>>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
@@ -64,88 +104,432 @@ impl ControlPlaneClient {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
|
||||
pub async fn heartbeat_ack(
|
||||
&self,
|
||||
authenticated_node_id: &str,
|
||||
authenticated_key: Option<&str>,
|
||||
authenticated_generation: &str,
|
||||
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::Http { .. } => {
|
||||
self.heartbeat_ack_http(
|
||||
authenticated_node_id,
|
||||
authenticated_key,
|
||||
authenticated_generation,
|
||||
payload,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControlPlaneMode::Local { .. } => {
|
||||
Err("local heartbeat callback requires connection credential binding".to_string())
|
||||
}
|
||||
ControlPlaneMode::Local { heartbeat_ack, .. } => heartbeat_ack(payload.to_vec()).await,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn heartbeat_ack_for_connection(
|
||||
&self,
|
||||
connection: Arc<ProxyConn>,
|
||||
payload: &[u8],
|
||||
) -> Result<Vec<u8>, String> {
|
||||
match self.inner.as_ref() {
|
||||
ControlPlaneMode::Disabled => Ok(b"{}".to_vec()),
|
||||
ControlPlaneMode::Http { .. } => {
|
||||
self.heartbeat_ack_http(
|
||||
&connection.node_id,
|
||||
connection.authenticated_key.as_deref(),
|
||||
&connection.node_generation,
|
||||
payload,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControlPlaneMode::Local { heartbeat_ack, .. } => {
|
||||
heartbeat_ack(connection, payload.to_vec()).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn heartbeat_ack_http(
|
||||
&self,
|
||||
authenticated_node_id: &str,
|
||||
authenticated_key: Option<&str>,
|
||||
authenticated_generation: &str,
|
||||
payload: &[u8],
|
||||
) -> Result<Vec<u8>, String> {
|
||||
let ControlPlaneMode::Http { client, base_url } = self.inner.as_ref() else {
|
||||
return Err("HTTP heartbeat callback is unavailable".to_string());
|
||||
};
|
||||
let Some(client) = client else {
|
||||
return Err("heartbeat callback HTTP client is unavailable".to_string());
|
||||
};
|
||||
let authenticated_key = authenticated_key
|
||||
.ok_or_else(|| "heartbeat callback is missing authenticated tunnel key".to_string())?;
|
||||
let url = format!("{}{TUNNEL_HEARTBEAT_PATH}", base_url.trim_end_matches('/'));
|
||||
let request = client
|
||||
.post(&url)
|
||||
.header("content-type", "application/json")
|
||||
.body(payload.to_vec());
|
||||
let response = sign_control_plane_request(
|
||||
request,
|
||||
authenticated_key,
|
||||
TUNNEL_HEARTBEAT_PATH,
|
||||
authenticated_node_id,
|
||||
authenticated_generation,
|
||||
payload,
|
||||
)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| {
|
||||
format!(
|
||||
"heartbeat callback request failed ({})",
|
||||
control_plane_reqwest_error_kind(&error)
|
||||
)
|
||||
})?;
|
||||
if response.status() == reqwest::StatusCode::FORBIDDEN {
|
||||
return Err(CONTROL_PLANE_CREDENTIAL_REVOKED.to_string());
|
||||
}
|
||||
if !response.status().is_success() {
|
||||
return Err(format!(
|
||||
"heartbeat callback failed with status {}",
|
||||
response.status()
|
||||
));
|
||||
}
|
||||
aether_http::read_response_bytes_with_limit(response, MAX_CONTROL_PLANE_RESPONSE_BYTES)
|
||||
.await
|
||||
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
|
||||
}
|
||||
|
||||
pub async fn push_node_status(
|
||||
&self,
|
||||
node_id: &str,
|
||||
authenticated_key: Option<&str>,
|
||||
authenticated_generation: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
) -> 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,
|
||||
"observed_at_unix_secs": observed_at_unix_secs,
|
||||
}))
|
||||
.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(),
|
||||
ControlPlaneMode::Http { .. } => {
|
||||
self.push_node_status_http(
|
||||
node_id,
|
||||
authenticated_key,
|
||||
authenticated_generation,
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControlPlaneMode::Local { .. } => {
|
||||
Err("local node-status callback requires connection credential binding".to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn push_node_status_for_connection(
|
||||
&self,
|
||||
connection: Arc<ProxyConn>,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
) -> Result<(), String> {
|
||||
match self.inner.as_ref() {
|
||||
ControlPlaneMode::Disabled => Ok(()),
|
||||
ControlPlaneMode::Http { .. } => {
|
||||
self.push_node_status_http(
|
||||
&connection.node_id,
|
||||
connection.authenticated_key.as_deref(),
|
||||
&connection.node_generation,
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControlPlaneMode::Local {
|
||||
push_node_status, ..
|
||||
} => push_node_status(connection, connected, conn_count, observed_at_unix_secs).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn push_node_status_http(
|
||||
&self,
|
||||
node_id: &str,
|
||||
authenticated_key: Option<&str>,
|
||||
authenticated_generation: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
) -> Result<(), String> {
|
||||
let ControlPlaneMode::Http { client, base_url } = self.inner.as_ref() else {
|
||||
return Err("HTTP node-status callback is unavailable".to_string());
|
||||
};
|
||||
let Some(client) = client else {
|
||||
return Err("node-status callback HTTP client is unavailable".to_string());
|
||||
};
|
||||
let authenticated_key = authenticated_key.ok_or_else(|| {
|
||||
"node-status callback is missing authenticated tunnel key".to_string()
|
||||
})?;
|
||||
let url = format!(
|
||||
"{}{TUNNEL_NODE_STATUS_PATH}",
|
||||
base_url.trim_end_matches('/')
|
||||
);
|
||||
let payload = serde_json::to_vec(&serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"connected": connected,
|
||||
"conn_count": conn_count,
|
||||
"observed_at_unix_secs": observed_at_unix_secs,
|
||||
}))
|
||||
.map_err(|e| format!("node-status callback serialization failed: {e}"))?;
|
||||
let request = client
|
||||
.post(&url)
|
||||
.header("content-type", "application/json")
|
||||
.body(payload.clone());
|
||||
let response = sign_control_plane_request(
|
||||
request,
|
||||
authenticated_key,
|
||||
TUNNEL_NODE_STATUS_PATH,
|
||||
node_id,
|
||||
authenticated_generation,
|
||||
&payload,
|
||||
)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| {
|
||||
format!(
|
||||
"node-status callback request failed ({})",
|
||||
control_plane_reqwest_error_kind(&error)
|
||||
)
|
||||
})?;
|
||||
if response.status() == reqwest::StatusCode::FORBIDDEN {
|
||||
return Err(CONTROL_PLANE_CREDENTIAL_REVOKED.to_string());
|
||||
}
|
||||
if response.status().is_success() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!(
|
||||
"node-status callback failed with status {}",
|
||||
response.status()
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_control_plane_base_url(raw: &str) -> Option<String> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty()
|
||||
|| raw.len() > MAX_CONTROL_PLANE_BASE_URL_BYTES
|
||||
|| raw.bytes().any(|byte| byte == 0 || byte.is_ascii_control())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
let parsed = url::Url::parse(raw).ok()?;
|
||||
if !matches!(parsed.scheme(), "http" | "https")
|
||||
|| parsed.host_str().is_none()
|
||||
|| !parsed.username().is_empty()
|
||||
|| parsed.password().is_some()
|
||||
|| parsed.query().is_some()
|
||||
|| parsed.fragment().is_some()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(raw.trim_end_matches('/').to_string())
|
||||
}
|
||||
|
||||
pub(crate) fn is_credential_revoked_error(error: &str) -> bool {
|
||||
error == CONTROL_PLANE_CREDENTIAL_REVOKED
|
||||
}
|
||||
|
||||
/// Return a stable transport category without rendering reqwest's error.
|
||||
///
|
||||
/// `reqwest::Error`'s `Display` implementation may include the complete URL
|
||||
/// (including path/query components). Control-plane errors are logged by the
|
||||
/// tunnel hub, so forwarding that value could disclose operator deployment
|
||||
/// details or credentials embedded in a path. Callers should use this helper
|
||||
/// whenever a control-plane request fails.
|
||||
fn control_plane_reqwest_error_kind(error: &reqwest::Error) -> &'static str {
|
||||
if error.is_timeout() {
|
||||
"timeout"
|
||||
} else if error.is_connect() {
|
||||
"connect"
|
||||
} else if error.is_request() {
|
||||
"request"
|
||||
} else if error.is_body() {
|
||||
"body"
|
||||
} else if error.is_decode() {
|
||||
"decode"
|
||||
} else {
|
||||
"transport"
|
||||
}
|
||||
}
|
||||
|
||||
fn sign_control_plane_request(
|
||||
request: reqwest::RequestBuilder,
|
||||
authenticated_key: &str,
|
||||
path: &str,
|
||||
node_id: &str,
|
||||
tunnel_generation: &str,
|
||||
body: &[u8],
|
||||
) -> Result<reqwest::RequestBuilder, String> {
|
||||
let timestamp = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map_err(|_| "system clock is before the Unix epoch".to_string())?
|
||||
.as_secs();
|
||||
let nonce = uuid::Uuid::new_v4().simple().to_string();
|
||||
let signature = sign_tunnel_control_plane_request_for_generation(
|
||||
authenticated_key,
|
||||
"POST",
|
||||
path,
|
||||
node_id,
|
||||
tunnel_generation,
|
||||
timestamp,
|
||||
&nonce,
|
||||
body,
|
||||
)
|
||||
.map_err(|error| format!("invalid authenticated tunnel key: {error}"))?;
|
||||
Ok(request
|
||||
.header(TUNNEL_CONTROL_PLANE_NODE_ID_HEADER, node_id)
|
||||
.header(TUNNEL_CONTROL_PLANE_GENERATION_HEADER, tunnel_generation)
|
||||
.header(TUNNEL_CONTROL_PLANE_TIMESTAMP_HEADER, timestamp)
|
||||
.header(TUNNEL_CONTROL_PLANE_NONCE_HEADER, nonce)
|
||||
.header(TUNNEL_CONTROL_PLANE_SIGNATURE_HEADER, signature))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::{
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
Arc,
|
||||
};
|
||||
|
||||
use axum::{
|
||||
body::Body,
|
||||
http::{header, Response, StatusCode},
|
||||
routing::post,
|
||||
Router,
|
||||
};
|
||||
use base64::Engine as _;
|
||||
|
||||
use super::{
|
||||
control_plane_reqwest_error_kind, normalize_control_plane_base_url, ControlPlaneClient,
|
||||
};
|
||||
use aether_gateway_tunnel::{TUNNEL_HEARTBEAT_PATH, TUNNEL_NODE_STATUS_PATH};
|
||||
|
||||
#[test]
|
||||
fn control_plane_base_url_rejects_credential_and_request_components() {
|
||||
for value in [
|
||||
"https://user:[email protected]",
|
||||
"https://gateway.example?token=secret",
|
||||
"https://gateway.example/control#fragment",
|
||||
"file:///tmp/gateway",
|
||||
"",
|
||||
] {
|
||||
assert!(
|
||||
normalize_control_plane_base_url(value).is_none(),
|
||||
"unsafe control-plane URL should be rejected: {value:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn control_plane_base_url_preserves_deployment_path_and_trims_slashes() {
|
||||
assert_eq!(
|
||||
normalize_control_plane_base_url(" https://gateway.example/control/// "),
|
||||
Some("https://gateway.example/control".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_control_plane_base_url("http://127.0.0.1:8084/"),
|
||||
Some("http://127.0.0.1:8084".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn control_plane_transport_errors_do_not_render_request_url() {
|
||||
let error = reqwest::Client::new()
|
||||
.get("ftp://user:[email protected]/control-plane")
|
||||
.send()
|
||||
.await
|
||||
.expect_err("unsupported control-plane scheme should fail before a request");
|
||||
let rendered = format!(
|
||||
"control-plane request failed ({})",
|
||||
control_plane_reqwest_error_kind(&error)
|
||||
);
|
||||
assert!(!rendered.contains("secret"));
|
||||
assert!(!rendered.contains("example.invalid"));
|
||||
assert!(rendered.starts_with("control-plane request failed ("));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn signed_control_plane_requests_never_follow_redirects() {
|
||||
let redirected_hits = Arc::new(AtomicUsize::new(0));
|
||||
let redirected_hits_for_route = Arc::clone(&redirected_hits);
|
||||
let redirected_listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("redirect target listener should bind");
|
||||
let redirected_addr = redirected_listener
|
||||
.local_addr()
|
||||
.expect("redirect target address should resolve");
|
||||
let redirected_app = Router::new().fallback(move || {
|
||||
let hits = Arc::clone(&redirected_hits_for_route);
|
||||
async move {
|
||||
hits.fetch_add(1, Ordering::SeqCst);
|
||||
StatusCode::OK
|
||||
}
|
||||
});
|
||||
let redirected_server = tokio::spawn(async move {
|
||||
axum::serve(redirected_listener, redirected_app)
|
||||
.await
|
||||
.expect("redirect target server should run");
|
||||
});
|
||||
|
||||
let source_listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("redirect source listener should bind");
|
||||
let source_addr = source_listener
|
||||
.local_addr()
|
||||
.expect("redirect source address should resolve");
|
||||
let location = format!("http://{redirected_addr}/captured");
|
||||
let redirect = move || {
|
||||
let location = location.clone();
|
||||
async move {
|
||||
Response::builder()
|
||||
.status(StatusCode::TEMPORARY_REDIRECT)
|
||||
.header(header::LOCATION, location)
|
||||
.body(Body::empty())
|
||||
.expect("redirect response should build")
|
||||
}
|
||||
};
|
||||
let source_app = Router::new()
|
||||
.route(TUNNEL_HEARTBEAT_PATH, post(redirect.clone()))
|
||||
.route(TUNNEL_NODE_STATUS_PATH, post(redirect));
|
||||
let source_server = tokio::spawn(async move {
|
||||
axum::serve(source_listener, source_app)
|
||||
.await
|
||||
.expect("redirect source server should run");
|
||||
});
|
||||
|
||||
let client = ControlPlaneClient::new(format!("http://{source_addr}"));
|
||||
let key = base64::engine::general_purpose::STANDARD.encode([7_u8; 32]);
|
||||
let heartbeat = client
|
||||
.heartbeat_ack(
|
||||
"node-1",
|
||||
Some(&key),
|
||||
"generation-1",
|
||||
br#"{"node_id":"node-1"}"#,
|
||||
)
|
||||
.await
|
||||
.expect_err("heartbeat redirect should be returned as an error");
|
||||
assert!(heartbeat.contains("307 Temporary Redirect"));
|
||||
let status = client
|
||||
.push_node_status("node-1", Some(&key), "generation-1", true, 1, 1)
|
||||
.await
|
||||
.expect_err("node-status redirect should be returned as an error");
|
||||
assert!(status.contains("307 Temporary Redirect"));
|
||||
assert_eq!(redirected_hits.load(Ordering::SeqCst), 0);
|
||||
|
||||
source_server.abort();
|
||||
redirected_server.abort();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::collections::{BTreeMap, HashMap, HashSet};
|
||||
use std::net::IpAddr;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicU8, AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, LazyLock};
|
||||
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
|
||||
@@ -19,6 +20,7 @@ use super::control_plane::ControlPlaneClient;
|
||||
use super::protocol;
|
||||
|
||||
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
|
||||
const MAX_TUNNEL_CONTROL_PAYLOAD_SIZE: usize = 256 * 1024;
|
||||
const SOFT_AVOID_QUEUE_PRESSURE_PERCENT: u64 = 50;
|
||||
const SOFT_AVOID_STREAM_PRESSURE_PERCENT: u64 = 85;
|
||||
const OUTBOUND_BACKPRESSURE_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
@@ -213,6 +215,9 @@ pub struct ProxyConn {
|
||||
pub id: u64,
|
||||
pub node_id: String,
|
||||
pub node_name: String,
|
||||
pub node_generation: String,
|
||||
pub authenticated_key: Option<String>,
|
||||
pub(crate) management_token_credential: Option<ProxyManagementTokenCredential>,
|
||||
pub outbound: BoundedOutbound,
|
||||
next_stream_id: AtomicU32,
|
||||
pub stream_count: AtomicUsize,
|
||||
@@ -242,6 +247,9 @@ impl ProxyConn {
|
||||
id,
|
||||
node_id,
|
||||
node_name,
|
||||
node_generation: String::new(),
|
||||
authenticated_key: None,
|
||||
management_token_credential: None,
|
||||
outbound: BoundedOutbound::new(tx, close_tx),
|
||||
next_stream_id: AtomicU32::new(2),
|
||||
stream_count: AtomicUsize::new(0),
|
||||
@@ -258,6 +266,37 @@ impl ProxyConn {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_authenticated_key(mut self, authenticated_key: String) -> Self {
|
||||
self.authenticated_key = Some(authenticated_key);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||
self.node_generation = tunnel_generation;
|
||||
self
|
||||
}
|
||||
|
||||
pub(crate) fn with_management_token_credential(
|
||||
mut self,
|
||||
credential: ProxyManagementTokenCredential,
|
||||
) -> Self {
|
||||
self.management_token_credential = Some(credential);
|
||||
self
|
||||
}
|
||||
|
||||
pub(crate) fn credential_binding(&self) -> Option<ProxyCredentialBinding> {
|
||||
match (
|
||||
self.authenticated_key.as_deref(),
|
||||
self.management_token_credential.as_ref(),
|
||||
) {
|
||||
(Some(key), None) => Some(ProxyCredentialBinding::Psk(key.to_string())),
|
||||
(None, Some(credential)) => {
|
||||
Some(ProxyCredentialBinding::ManagementToken(credential.clone()))
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn record_write_latency(&self, elapsed: std::time::Duration) {
|
||||
let micros = u64::try_from(elapsed.as_micros()).unwrap_or(u64::MAX);
|
||||
self.write_latency_last_us.store(micros, Ordering::Relaxed);
|
||||
@@ -440,6 +479,44 @@ impl ProxyConn {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct ProxyManagementTokenCredential {
|
||||
pub(crate) verified_token_hash: crate::management_token_auth::VerifiedManagementTokenHash,
|
||||
pub(crate) token_id: String,
|
||||
pub(crate) user_id: String,
|
||||
pub(crate) remote_ip: IpAddr,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) enum ProxyCredentialBinding {
|
||||
Psk(String),
|
||||
ManagementToken(ProxyManagementTokenCredential),
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProxyCredentialBinding {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Psk(_) => formatter.write_str("ProxyCredentialBinding::Psk([REDACTED])"),
|
||||
Self::ManagementToken(credential) => formatter
|
||||
.debug_tuple("ProxyCredentialBinding::ManagementToken")
|
||||
.field(credential)
|
||||
.finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProxyManagementTokenCredential {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProxyManagementTokenCredential")
|
||||
.field("verified_token_hash", &"[REDACTED]")
|
||||
.field("token_id", &self.token_id)
|
||||
.field("user_id", &self.user_id)
|
||||
.field("remote_ip", &self.remote_ip)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct ProxyConnSnapshot {
|
||||
conn_id: u64,
|
||||
@@ -503,6 +580,7 @@ struct LocalWaitState {
|
||||
|
||||
pub struct LocalStream {
|
||||
pub id: u64,
|
||||
tunnel_generation: String,
|
||||
proxy_conn_id: u64,
|
||||
proxy_stream_id: u32,
|
||||
request_window: StreamFlowWindow,
|
||||
@@ -515,10 +593,17 @@ pub struct LocalStream {
|
||||
}
|
||||
|
||||
impl LocalStream {
|
||||
fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32, initial_window_bytes: u32) -> Self {
|
||||
fn new(
|
||||
id: u64,
|
||||
tunnel_generation: String,
|
||||
proxy_conn_id: u64,
|
||||
proxy_stream_id: u32,
|
||||
initial_window_bytes: u32,
|
||||
) -> Self {
|
||||
let (body_tx, body_rx) = mpsc::channel(128);
|
||||
Self {
|
||||
id,
|
||||
tunnel_generation,
|
||||
proxy_conn_id,
|
||||
proxy_stream_id,
|
||||
request_window: StreamFlowWindow::new(initial_window_bytes),
|
||||
@@ -531,6 +616,10 @@ impl LocalStream {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn tunnel_generation(&self) -> &str {
|
||||
&self.tunnel_generation
|
||||
}
|
||||
|
||||
async fn acquire_request_window(
|
||||
&self,
|
||||
bytes: usize,
|
||||
@@ -676,9 +765,11 @@ pub struct HubRouter {
|
||||
drain_reasons: Mutex<HashMap<String, u64>>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct NodeStatusEvent {
|
||||
node_id: String,
|
||||
authenticated_key: Option<String>,
|
||||
tunnel_generation: String,
|
||||
connection: Option<Arc<ProxyConn>>,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
observed_at_unix_secs: u64,
|
||||
@@ -692,15 +783,37 @@ impl HubRouter {
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
handle.spawn(async move {
|
||||
while let Some(event) = node_status_rx.recv().await {
|
||||
if let Err(error) = worker_control_plane
|
||||
.push_node_status(
|
||||
&event.node_id,
|
||||
event.connected,
|
||||
event.conn_count,
|
||||
event.observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
{
|
||||
let connection = event.connection.clone();
|
||||
let result = match connection {
|
||||
Some(connection) => {
|
||||
worker_control_plane
|
||||
.push_node_status_for_connection(
|
||||
connection,
|
||||
event.connected,
|
||||
event.conn_count,
|
||||
event.observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
None => {
|
||||
worker_control_plane
|
||||
.push_node_status(
|
||||
&event.node_id,
|
||||
event.authenticated_key.as_deref(),
|
||||
&event.tunnel_generation,
|
||||
event.connected,
|
||||
event.conn_count,
|
||||
event.observed_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
};
|
||||
if let Err(error) = result {
|
||||
if super::control_plane::is_credential_revoked_error(&error) {
|
||||
if let Some(connection) = event.connection.as_ref() {
|
||||
connection.request_close();
|
||||
}
|
||||
}
|
||||
warn!(
|
||||
node_id = %event.node_id,
|
||||
connected = event.connected,
|
||||
@@ -746,7 +859,7 @@ impl HubRouter {
|
||||
|
||||
let healthy_count = {
|
||||
let mut map = self.proxy_conns.write();
|
||||
map.entry(node_id.clone()).or_default().push(conn);
|
||||
map.entry(node_id.clone()).or_default().push(conn.clone());
|
||||
available_conn_count(map.get(&node_id).map(Vec::as_slice).unwrap_or(&[]))
|
||||
};
|
||||
|
||||
@@ -758,11 +871,23 @@ impl HubRouter {
|
||||
"proxy connected"
|
||||
);
|
||||
|
||||
self.notify_node_status(node_id, healthy_count > 0, healthy_count);
|
||||
self.notify_node_status(
|
||||
node_id,
|
||||
conn.authenticated_key.clone(),
|
||||
Some(conn),
|
||||
healthy_count > 0,
|
||||
healthy_count,
|
||||
);
|
||||
}
|
||||
|
||||
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
|
||||
self.proxy_conns_by_id.remove(&conn_id);
|
||||
let disconnected_connection = self
|
||||
.proxy_conns_by_id
|
||||
.remove(&conn_id)
|
||||
.map(|(_, connection)| connection);
|
||||
let disconnected_authenticated_key = disconnected_connection
|
||||
.as_ref()
|
||||
.and_then(|connection| connection.authenticated_key.clone());
|
||||
|
||||
let healthy_count = {
|
||||
let mut map = self.proxy_conns.write();
|
||||
@@ -785,7 +910,20 @@ impl HubRouter {
|
||||
);
|
||||
|
||||
self.cancel_streams_for_proxy(conn_id);
|
||||
self.notify_node_status(node_id.to_string(), healthy_count > 0, healthy_count);
|
||||
let authenticated_key = self
|
||||
.proxy_conns
|
||||
.read()
|
||||
.get(node_id)
|
||||
.and_then(|connections| connections.first())
|
||||
.and_then(|connection| connection.authenticated_key.clone())
|
||||
.or(disconnected_authenticated_key);
|
||||
self.notify_node_status(
|
||||
node_id.to_string(),
|
||||
authenticated_key,
|
||||
disconnected_connection,
|
||||
healthy_count > 0,
|
||||
healthy_count,
|
||||
);
|
||||
}
|
||||
|
||||
pub fn request_close_all_proxies(&self) -> usize {
|
||||
@@ -801,9 +939,52 @@ impl HubRouter {
|
||||
total
|
||||
}
|
||||
|
||||
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
|
||||
pub(crate) fn request_close_proxy(&self, conn_id: u64) -> bool {
|
||||
let Some(conn) = self
|
||||
.proxy_conns_by_id
|
||||
.get(&conn_id)
|
||||
.map(|entry| Arc::clone(entry.value()))
|
||||
else {
|
||||
return false;
|
||||
};
|
||||
conn.request_close();
|
||||
true
|
||||
}
|
||||
|
||||
pub(crate) fn request_close_proxies_for_node(&self, node_id: &str) -> usize {
|
||||
let conns = self.proxy_connections_for_node(node_id);
|
||||
let total = conns.len();
|
||||
for conn in conns {
|
||||
conn.request_close();
|
||||
}
|
||||
total
|
||||
}
|
||||
|
||||
pub(crate) fn proxy_connections_for_node(&self, node_id: &str) -> Vec<Arc<ProxyConn>> {
|
||||
self.proxy_conns
|
||||
.read()
|
||||
.get(node_id)
|
||||
.map(|connections| connections.to_vec())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn notify_node_status(
|
||||
&self,
|
||||
node_id: String,
|
||||
authenticated_key: Option<String>,
|
||||
connection: Option<Arc<ProxyConn>>,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
) {
|
||||
let tunnel_generation = connection
|
||||
.as_ref()
|
||||
.map(|connection| connection.node_generation.clone())
|
||||
.unwrap_or_default();
|
||||
let event = NodeStatusEvent {
|
||||
node_id,
|
||||
authenticated_key,
|
||||
tunnel_generation,
|
||||
connection,
|
||||
connected,
|
||||
conn_count,
|
||||
observed_at_unix_secs: current_unix_secs(),
|
||||
@@ -837,13 +1018,24 @@ impl HubRouter {
|
||||
}
|
||||
|
||||
fn notify_current_node_status(&self, node_id: &str) {
|
||||
let healthy_count = {
|
||||
let (healthy_count, connection, authenticated_key) = {
|
||||
let map = self.proxy_conns.read();
|
||||
map.get(node_id)
|
||||
.map(|v| available_conn_count(v.as_slice()))
|
||||
.unwrap_or(0)
|
||||
let connections = map.get(node_id).map(Vec::as_slice).unwrap_or(&[]);
|
||||
(
|
||||
available_conn_count(connections),
|
||||
connections.first().cloned(),
|
||||
connections
|
||||
.first()
|
||||
.and_then(|connection| connection.authenticated_key.clone()),
|
||||
)
|
||||
};
|
||||
self.notify_node_status(node_id.to_string(), healthy_count > 0, healthy_count);
|
||||
self.notify_node_status(
|
||||
node_id.to_string(),
|
||||
authenticated_key,
|
||||
connection,
|
||||
healthy_count > 0,
|
||||
healthy_count,
|
||||
);
|
||||
}
|
||||
|
||||
fn record_stream_reset(&self, reason: &str) {
|
||||
@@ -856,7 +1048,11 @@ impl HubRouter {
|
||||
increment_reason(&self.drain_reasons, reason);
|
||||
}
|
||||
|
||||
fn ranked_proxy_conn_candidates(&self, node_id: &str) -> Vec<ProxyConnCandidate> {
|
||||
fn ranked_proxy_conn_candidates(
|
||||
&self,
|
||||
node_id: &str,
|
||||
authorized_conn_ids: Option<&HashSet<u64>>,
|
||||
) -> Vec<ProxyConnCandidate> {
|
||||
let conns = {
|
||||
let map = self.proxy_conns.read();
|
||||
map.get(node_id)
|
||||
@@ -866,6 +1062,9 @@ impl HubRouter {
|
||||
let mut candidates = conns
|
||||
.into_iter()
|
||||
.filter_map(|conn| {
|
||||
if authorized_conn_ids.is_some_and(|allowed| !allowed.contains(&conn.id)) {
|
||||
return None;
|
||||
}
|
||||
let snapshot = conn.snapshot();
|
||||
snapshot
|
||||
.available
|
||||
@@ -877,7 +1076,7 @@ impl HubRouter {
|
||||
}
|
||||
|
||||
pub fn has_local_proxy(&self, node_id: &str) -> bool {
|
||||
!self.ranked_proxy_conn_candidates(node_id).is_empty()
|
||||
!self.ranked_proxy_conn_candidates(node_id, None).is_empty()
|
||||
}
|
||||
|
||||
pub async fn open_local_stream(
|
||||
@@ -885,7 +1084,17 @@ impl HubRouter {
|
||||
node_id: &str,
|
||||
meta: &protocol::RequestMeta,
|
||||
) -> Result<Arc<LocalStream>, String> {
|
||||
let candidates = self.ranked_proxy_conn_candidates(node_id);
|
||||
self.open_local_stream_with_authorized_connections(node_id, meta, None)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn open_local_stream_with_authorized_connections(
|
||||
&self,
|
||||
node_id: &str,
|
||||
meta: &protocol::RequestMeta,
|
||||
authorized_conn_ids: Option<&HashSet<u64>>,
|
||||
) -> Result<Arc<LocalStream>, String> {
|
||||
let candidates = self.ranked_proxy_conn_candidates(node_id, authorized_conn_ids);
|
||||
if candidates.is_empty() {
|
||||
self.selection_unavailable_total
|
||||
.fetch_add(1, Ordering::Relaxed);
|
||||
@@ -959,6 +1168,7 @@ impl HubRouter {
|
||||
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.node_generation.clone(),
|
||||
proxy_conn.id,
|
||||
proxy_stream_id,
|
||||
*STREAM_INITIAL_WINDOW_BYTES,
|
||||
@@ -1160,8 +1370,14 @@ impl HubRouter {
|
||||
Some(h) => h,
|
||||
None => return,
|
||||
};
|
||||
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
|
||||
if data.len() < expected_len {
|
||||
let Some(expected_len) = protocol::HEADER_SIZE.checked_add(header.payload_len as usize)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if data.len() != expected_len {
|
||||
if header.stream_id != 0 {
|
||||
self.fail_proxy_stream(proxy_conn_id, header.stream_id, "invalid frame length");
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -1176,22 +1392,32 @@ impl HubRouter {
|
||||
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());
|
||||
let raw_message = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| String::from_utf8(payload).ok())
|
||||
.unwrap_or_else(|| "stream error".to_string());
|
||||
let message = safe_peer_stream_error(&raw_message);
|
||||
self.record_stream_reset(&message);
|
||||
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
|
||||
}
|
||||
protocol::RESET_STREAM => {
|
||||
let message = protocol::decode_payload(data, &header)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::ResetStreamPayload>(&payload)
|
||||
.ok()
|
||||
.map(|payload| payload.reason)
|
||||
})
|
||||
.unwrap_or_else(|| "stream reset".to_string());
|
||||
let raw_message = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::ResetStreamPayload>(&payload)
|
||||
.ok()
|
||||
.map(|payload| payload.reason)
|
||||
})
|
||||
.unwrap_or_else(|| "stream reset".to_string());
|
||||
let message = safe_peer_stream_error(&raw_message);
|
||||
self.record_stream_reset(&message);
|
||||
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
|
||||
}
|
||||
@@ -1217,17 +1443,19 @@ impl HubRouter {
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
let first = pc.mark_draining();
|
||||
if first {
|
||||
let drain =
|
||||
protocol::decode_payload(data, &header)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
if payload.is_empty() {
|
||||
None
|
||||
} else {
|
||||
serde_json::from_slice::<protocol::GoAwayPayload>(&payload)
|
||||
.ok()
|
||||
}
|
||||
});
|
||||
let drain = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
if payload.is_empty() {
|
||||
None
|
||||
} else {
|
||||
serde_json::from_slice::<protocol::GoAwayPayload>(&payload).ok()
|
||||
}
|
||||
});
|
||||
let reason = drain
|
||||
.as_ref()
|
||||
.map(|payload| payload.reason.as_str())
|
||||
@@ -1262,12 +1490,13 @@ impl HubRouter {
|
||||
}
|
||||
}
|
||||
protocol::HELLO => {
|
||||
if let Some(payload) =
|
||||
protocol::decode_payload(data, &header)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::HelloPayload>(&payload).ok()
|
||||
})
|
||||
if let Some(payload) = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| serde_json::from_slice::<protocol::HelloPayload>(&payload).ok())
|
||||
{
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
pc.update_protocol_version(payload.protocol_version);
|
||||
@@ -1290,13 +1519,15 @@ impl HubRouter {
|
||||
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
||||
}
|
||||
protocol::LOAD_REPORT => {
|
||||
if let Some(payload) =
|
||||
protocol::decode_payload(data, &header)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::LoadReportPayload>(&payload).ok()
|
||||
})
|
||||
{
|
||||
if let Some(payload) = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::LoadReportPayload>(&payload).ok()
|
||||
}) {
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
pc.update_remote_health_score(payload.health_score);
|
||||
}
|
||||
@@ -1332,12 +1563,13 @@ impl HubRouter {
|
||||
data: &[u8],
|
||||
header: &protocol::FrameHeader,
|
||||
) {
|
||||
let Some(delta) = protocol::decode_payload(data, header)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::WindowUpdatePayload>(&payload).ok()
|
||||
})
|
||||
.map(|payload| payload.delta_bytes)
|
||||
let Some(delta) =
|
||||
protocol::decode_payload_with_limit(data, header, MAX_TUNNEL_CONTROL_PAYLOAD_SIZE)
|
||||
.ok()
|
||||
.and_then(|payload| {
|
||||
serde_json::from_slice::<protocol::WindowUpdatePayload>(&payload).ok()
|
||||
})
|
||||
.map(|payload| payload.delta_bytes)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
@@ -1457,7 +1689,9 @@ impl HubRouter {
|
||||
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 {
|
||||
let Ok(payload) =
|
||||
protocol::decode_payload_with_limit(data, &header, MAX_TUNNEL_CONTROL_PAYLOAD_SIZE)
|
||||
else {
|
||||
self.fail_proxy_stream(
|
||||
proxy_conn_id,
|
||||
header.stream_id,
|
||||
@@ -1487,7 +1721,11 @@ impl HubRouter {
|
||||
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 {
|
||||
let Ok(payload) = protocol::decode_payload_with_limit(
|
||||
data,
|
||||
&header,
|
||||
aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES,
|
||||
) else {
|
||||
self.fail_proxy_stream(
|
||||
proxy_conn_id,
|
||||
header.stream_id,
|
||||
@@ -1567,16 +1805,70 @@ impl HubRouter {
|
||||
data: &[u8],
|
||||
header: &protocol::FrameHeader,
|
||||
) {
|
||||
let payload = match protocol::decode_payload(data, header) {
|
||||
let payload = match protocol::decode_payload_with_limit(
|
||||
data,
|
||||
header,
|
||||
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||
) {
|
||||
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 {
|
||||
let Some(authenticated_node_id) = self
|
||||
.proxy_conns_by_id
|
||||
.get(&proxy_conn_id)
|
||||
.map(|entry| entry.node_id.clone())
|
||||
else {
|
||||
warn!(
|
||||
proxy_conn_id,
|
||||
"heartbeat rejected for unregistered proxy connection"
|
||||
);
|
||||
return;
|
||||
};
|
||||
let payload_node_id = match heartbeat_payload_node_id(&payload) {
|
||||
Ok(node_id) => node_id,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
proxy_conn_id,
|
||||
authenticated_node_id = %authenticated_node_id,
|
||||
error = %error,
|
||||
"heartbeat rejected before control-plane dispatch"
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if payload_node_id != authenticated_node_id {
|
||||
warn!(
|
||||
proxy_conn_id,
|
||||
authenticated_node_id = %authenticated_node_id,
|
||||
payload_node_id = %payload_node_id,
|
||||
"heartbeat rejected because node identity does not match tunnel authentication"
|
||||
);
|
||||
return;
|
||||
}
|
||||
let connection = self
|
||||
.proxy_conns_by_id
|
||||
.get(&proxy_conn_id)
|
||||
.map(|entry| Arc::clone(entry.value()));
|
||||
let Some(connection) = connection else {
|
||||
warn!(
|
||||
proxy_conn_id,
|
||||
"heartbeat rejected for unregistered proxy connection"
|
||||
);
|
||||
return;
|
||||
};
|
||||
let ack_payload = match self
|
||||
.control_plane
|
||||
.heartbeat_ack_for_connection(connection.clone(), &payload)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
if super::control_plane::is_credential_revoked_error(&error) {
|
||||
connection.request_close();
|
||||
}
|
||||
warn!(
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
error = %error,
|
||||
@@ -1755,6 +2047,18 @@ impl HubRouter {
|
||||
}
|
||||
}
|
||||
|
||||
fn heartbeat_payload_node_id(payload: &[u8]) -> Result<String, String> {
|
||||
let value: serde_json::Value = serde_json::from_slice(payload)
|
||||
.map_err(|_| "heartbeat payload is not valid JSON".to_string())?;
|
||||
let node_id = value
|
||||
.get("node_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "heartbeat payload is missing node_id".to_string())?;
|
||||
Ok(node_id.to_string())
|
||||
}
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -1830,6 +2134,56 @@ fn metric_reason(raw: &str) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
fn safe_peer_stream_error(raw: &str) -> String {
|
||||
const CLASSIFICATION_PREFIX_BYTES: usize = 4 * 1024;
|
||||
|
||||
let prefix = if raw.len() <= CLASSIFICATION_PREFIX_BYTES {
|
||||
raw
|
||||
} else {
|
||||
let mut end = CLASSIFICATION_PREFIX_BYTES;
|
||||
while !raw.is_char_boundary(end) {
|
||||
end = end.saturating_sub(1);
|
||||
}
|
||||
&raw[..end]
|
||||
};
|
||||
let category = if contains_ascii_case_insensitive(prefix, "timed out")
|
||||
|| contains_ascii_case_insensitive(prefix, "timeout")
|
||||
{
|
||||
"timeout"
|
||||
} else if contains_ascii_case_insensitive(prefix, "overloaded")
|
||||
|| contains_ascii_case_insensitive(prefix, "backpressure")
|
||||
|| contains_ascii_case_insensitive(prefix, "congested")
|
||||
|| contains_ascii_case_insensitive(prefix, "window")
|
||||
{
|
||||
"overloaded"
|
||||
} else if contains_ascii_case_insensitive(prefix, "forbidden")
|
||||
|| contains_ascii_case_insensitive(prefix, "unauthorized")
|
||||
|| contains_ascii_case_insensitive(prefix, "authentication")
|
||||
{
|
||||
"forbidden"
|
||||
} else if contains_ascii_case_insensitive(prefix, "dns") {
|
||||
"dns"
|
||||
} else if contains_ascii_case_insensitive(prefix, "connect")
|
||||
|| contains_ascii_case_insensitive(prefix, "socket")
|
||||
{
|
||||
"connect"
|
||||
} else if contains_ascii_case_insensitive(prefix, "cancel")
|
||||
|| contains_ascii_case_insensitive(prefix, "reset")
|
||||
{
|
||||
"reset"
|
||||
} else {
|
||||
"relay"
|
||||
};
|
||||
format!("tunnel stream {category} error")
|
||||
}
|
||||
|
||||
fn contains_ascii_case_insensitive(haystack: &str, needle: &str) -> bool {
|
||||
haystack
|
||||
.as_bytes()
|
||||
.windows(needle.len())
|
||||
.any(|candidate| candidate.eq_ignore_ascii_case(needle.as_bytes()))
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct HubStats {
|
||||
pub proxy_connections: usize,
|
||||
@@ -2098,13 +2452,36 @@ impl HubStats {
|
||||
mod tests {
|
||||
use aether_runtime::bounded_queue;
|
||||
|
||||
use super::{protocol, ControlPlaneClient, HubRouter, ProxyConn, MAX_REQUEST_BODY_FRAME_SIZE};
|
||||
use super::{
|
||||
protocol, safe_peer_stream_error, ControlPlaneClient, HubRouter, ProxyConn,
|
||||
MAX_REQUEST_BODY_FRAME_SIZE,
|
||||
};
|
||||
use axum::extract::ws::Message;
|
||||
use bytes::Bytes;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::watch;
|
||||
|
||||
#[test]
|
||||
fn peer_stream_errors_are_projected_to_finite_categories() {
|
||||
let sensitive = safe_peer_stream_error(
|
||||
"request failed for https://alice:[email protected]/private?token=query-secret\r\nx: y",
|
||||
);
|
||||
assert_eq!(sensitive, "tunnel stream relay error");
|
||||
for secret in ["alice", "secret", "10.0.0.8", "private", "query", "x: y"] {
|
||||
assert!(!sensitive.contains(secret), "leaked {secret}: {sensitive}");
|
||||
}
|
||||
assert_eq!(
|
||||
safe_peer_stream_error("upstream connect timeout: Bearer secret"),
|
||||
"tunnel stream timeout error"
|
||||
);
|
||||
assert_eq!(
|
||||
safe_peer_stream_error("outbound backpressure timeout"),
|
||||
"tunnel stream timeout error"
|
||||
);
|
||||
}
|
||||
|
||||
fn build_meta() -> protocol::RequestMeta {
|
||||
protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
@@ -2358,8 +2735,10 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn heartbeat_callback_failure_does_not_send_fake_ack() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::local(
|
||||
|_payload| Box::pin(async { Err("db unavailable".to_string()) }),
|
||||
|_node_id, _connected, _conn_count, _observed_at_unix_secs| Box::pin(async { Ok(()) }),
|
||||
|_connection, _payload| Box::pin(async { Err("db unavailable".to_string()) }),
|
||||
|_connection, _connected, _conn_count, _observed_at_unix_secs| {
|
||||
Box::pin(async { Ok(()) })
|
||||
},
|
||||
));
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
|
||||
@@ -2386,6 +2765,45 @@ mod tests {
|
||||
assert!(proxy_rx.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn heartbeat_for_another_node_is_rejected_before_callback_and_ack() {
|
||||
let callback_calls = Arc::new(AtomicUsize::new(0));
|
||||
let callback_calls_for_heartbeat = Arc::clone(&callback_calls);
|
||||
let hub = HubRouter::new(ControlPlaneClient::local(
|
||||
move |_connection, _payload| {
|
||||
callback_calls_for_heartbeat.fetch_add(1, Ordering::Relaxed);
|
||||
Box::pin(async { Ok(br#"{"heartbeat_id":99}"#.to_vec()) })
|
||||
},
|
||||
|_connection, _connected, _conn_count, _observed_at_unix_secs| {
|
||||
Box::pin(async { Ok(()) })
|
||||
},
|
||||
));
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let proxy = Arc::new(ProxyConn::new(
|
||||
301,
|
||||
"authenticated-node".to_string(),
|
||||
"Authenticated Node".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
2,
|
||||
));
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let payload = serde_json::to_vec(&serde_json::json!({
|
||||
"node_id": "victim-node",
|
||||
"heartbeat_id": 99u64,
|
||||
}))
|
||||
.expect("payload should serialize");
|
||||
let mut frame = protocol::encode_frame(1, protocol::HEARTBEAT_DATA, 0, &payload);
|
||||
hub.handle_proxy_frame(301, &mut frame).await;
|
||||
|
||||
assert_eq!(callback_calls.load(Ordering::Relaxed), 0);
|
||||
assert!(proxy_rx.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn second_stream_works_after_first_completes_via_stream_end() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
|
||||
@@ -2,28 +2,23 @@ use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_contracts::tunnel::{
|
||||
resolve_tunnel_request_timeouts, try_decode_tunnel_relay_request_meta,
|
||||
TUNNEL_RELAY_FORWARDED_BY_HEADER,
|
||||
};
|
||||
use aether_contracts::tunnel::{resolve_tunnel_request_timeouts, 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 tokio::sync::mpsc;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::api::response::apply_streaming_response_headers;
|
||||
use crate::headers::should_skip_response_header;
|
||||
use crate::maintenance::record_proxy_upgrade_traffic_success;
|
||||
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
||||
|
||||
use super::hub::{LocalBodyEvent, LocalStream};
|
||||
use super::protocol;
|
||||
use super::AppState;
|
||||
use super::{AppState, RelayRequestAuthenticated};
|
||||
|
||||
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||
|
||||
@@ -86,8 +81,7 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
.await
|
||||
.map_err(map_request_admission_error)?;
|
||||
let stream = state
|
||||
.hub
|
||||
.open_local_stream(node_id, &meta)
|
||||
.open_authorized_local_stream(node_id, &meta)
|
||||
.await
|
||||
.map_err(|error| format!("connect: {error}"))?;
|
||||
if let Err(error) = state
|
||||
@@ -107,7 +101,13 @@ pub(crate) async fn open_direct_relay_stream(
|
||||
return Err(format!("timeout: {error}"));
|
||||
}
|
||||
};
|
||||
if let Err(error) = record_proxy_upgrade_traffic_success(state.data.as_ref(), node_id).await {
|
||||
if let Err(error) = record_proxy_upgrade_traffic_success_for_generation(
|
||||
state.data.as_ref(),
|
||||
node_id,
|
||||
stream.tunnel_generation(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
error = %error,
|
||||
@@ -171,7 +171,7 @@ fn is_rollout_probe_request(headers: &HeaderMap, forwarded_by_gateway: bool) ->
|
||||
pub async fn relay_request(
|
||||
Path(node_id): Path<String>,
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
||||
ConnectInfo(_addr): ConnectInfo<SocketAddr>,
|
||||
request: Request,
|
||||
) -> impl IntoResponse {
|
||||
let forwarded_by_gateway = request
|
||||
@@ -181,13 +181,10 @@ pub async fn relay_request(
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
let rollout_probe = is_rollout_probe_request(request.headers(), forwarded_by_gateway);
|
||||
if !addr.ip().is_loopback() && !forwarded_by_gateway {
|
||||
return tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
"forbidden",
|
||||
"local relay only accepts loopback requests",
|
||||
);
|
||||
}
|
||||
let already_authenticated = request
|
||||
.extensions()
|
||||
.get::<RelayRequestAuthenticated>()
|
||||
.is_some();
|
||||
|
||||
let request_permit = match state.try_acquire_request_permit().await {
|
||||
Ok(permit) => permit,
|
||||
@@ -226,109 +223,77 @@ pub async fn relay_request(
|
||||
}
|
||||
};
|
||||
|
||||
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;
|
||||
if !already_authenticated {
|
||||
return release_permit_response(
|
||||
tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
"forbidden",
|
||||
"relay request integrity must be verified before local dispatch",
|
||||
),
|
||||
request_permit,
|
||||
);
|
||||
}
|
||||
|
||||
while let Some(chunk_result) = body_stream.next().await {
|
||||
let chunk = match chunk_result {
|
||||
Ok(chunk) => chunk,
|
||||
let Some(spool) = request
|
||||
.extensions()
|
||||
.get::<crate::tunnel::VerifiedRelaySpool>()
|
||||
.cloned()
|
||||
else {
|
||||
return release_permit_response(
|
||||
tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
"forbidden",
|
||||
"verified relay payload is missing",
|
||||
),
|
||||
request_permit,
|
||||
);
|
||||
};
|
||||
let meta = spool.meta().clone();
|
||||
|
||||
let stream = match state.open_authorized_local_stream(&node_id, &meta).await {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
return release_permit_response(
|
||||
tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error),
|
||||
request_permit,
|
||||
);
|
||||
}
|
||||
};
|
||||
let body_stream = match spool.body_stream().await {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return release_permit_response(
|
||||
tunnel_error_response(StatusCode::BAD_GATEWAY, "relay", &error),
|
||||
request_permit,
|
||||
);
|
||||
}
|
||||
};
|
||||
futures_util::pin_mut!(body_stream);
|
||||
while let Some(chunk) = futures_util::StreamExt::next(&mut body_stream).await {
|
||||
let (chunk, end) = match chunk {
|
||||
Ok(chunk) => (chunk, false),
|
||||
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");
|
||||
let error = error.to_string();
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return release_permit_response(
|
||||
tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"failed to read relay request body",
|
||||
),
|
||||
tunnel_error_response(StatusCode::BAD_GATEWAY, "relay", &error),
|
||||
request_permit,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if stream.is_none() {
|
||||
envelope_buf.extend_from_slice(&chunk);
|
||||
let Some((parsed_meta, body_offset)) =
|
||||
(match try_decode_tunnel_relay_request_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).await {
|
||||
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)
|
||||
.await
|
||||
{
|
||||
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)
|
||||
.push_local_request_body(stream.id, chunk, end)
|
||||
.await
|
||||
{
|
||||
state.hub.cancel_local_stream(active_stream.id, &error);
|
||||
state.hub.cancel_local_stream(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)
|
||||
@@ -359,8 +324,12 @@ pub async fn relay_request(
|
||||
}
|
||||
};
|
||||
if !rollout_probe {
|
||||
if let Err(error) =
|
||||
record_proxy_upgrade_traffic_success(state.data.as_ref(), &node_id).await
|
||||
if let Err(error) = record_proxy_upgrade_traffic_success_for_generation(
|
||||
state.data.as_ref(),
|
||||
&node_id,
|
||||
stream.tunnel_generation(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
@@ -436,8 +405,16 @@ fn release_permit_response(
|
||||
}
|
||||
|
||||
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
|
||||
let connection_declared = aether_http::connection_declared_header_names(
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(name, _)| name.eq_ignore_ascii_case(http::header::CONNECTION.as_str()))
|
||||
.map(|(_, value)| value.as_str()),
|
||||
);
|
||||
for (name, value) in headers {
|
||||
if should_skip_local_relay_response_header(name) {
|
||||
if should_skip_local_relay_response_header(name)
|
||||
|| connection_declared.contains(&name.to_ascii_lowercase())
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
|
||||
@@ -454,12 +431,21 @@ fn should_skip_local_relay_response_header(name: &str) -> bool {
|
||||
should_skip_response_header(name) || name.eq_ignore_ascii_case("content-length")
|
||||
}
|
||||
|
||||
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
|
||||
fn tunnel_error_response(status: StatusCode, kind: &str, _message: &str) -> Response<Body> {
|
||||
let kind = safe_tunnel_error_kind(kind);
|
||||
let message = match kind {
|
||||
"overloaded" => "hub relay overloaded",
|
||||
"forbidden" => "relay request forbidden",
|
||||
"connect" => "tunnel connection failed",
|
||||
"timeout" => "tunnel request timed out",
|
||||
"unavailable" => "tunnel unavailable",
|
||||
_ => "tunnel relay failed",
|
||||
};
|
||||
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")),
|
||||
HeaderValue::from_static(kind),
|
||||
);
|
||||
headers.insert(
|
||||
axum::http::header::CONTENT_TYPE,
|
||||
@@ -471,17 +457,32 @@ fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Respo
|
||||
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
|
||||
}
|
||||
|
||||
fn safe_tunnel_error_kind(kind: &str) -> &'static str {
|
||||
match kind.trim().to_ascii_lowercase().as_str() {
|
||||
"overloaded" => "overloaded",
|
||||
"forbidden" => "forbidden",
|
||||
"connect" => "connect",
|
||||
"timeout" => "timeout",
|
||||
"unavailable" => "unavailable",
|
||||
"relay" => "relay",
|
||||
_ => "relay",
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::hub::ProxyConn;
|
||||
use super::super::{protocol, AppState, ConnConfig, ControlPlaneClient};
|
||||
use super::super::{
|
||||
protocol, AppState, ConnConfig, ControlPlaneClient, RelayRequestAuthenticated,
|
||||
};
|
||||
use super::{
|
||||
is_rollout_probe_request, relay_header_timeout, relay_request, Body, HeaderMap, Request,
|
||||
SocketAddr, StatusCode, TUNNEL_ERROR_HEADER,
|
||||
is_rollout_probe_request, relay_header_timeout, relay_request, tunnel_error_response, Body,
|
||||
HeaderMap, Request, SocketAddr, StatusCode, TUNNEL_ERROR_HEADER,
|
||||
};
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::maintenance::start_proxy_upgrade_rollout;
|
||||
use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER;
|
||||
use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED;
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
InMemoryProxyNodeRepository, ProxyNodeHeartbeatMutation, ProxyNodeWriteRepository,
|
||||
StoredProxyNode,
|
||||
@@ -496,6 +497,34 @@ mod tests {
|
||||
use std::time::Duration;
|
||||
use tokio::sync::watch;
|
||||
|
||||
const LOCAL_TUNNEL_TEST_PSK: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||
const LOCAL_TUNNEL_TEST_GENERATION: &str = "local-relay-test-generation-1";
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_error_response_drops_internal_and_peer_details() {
|
||||
let response = tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"https://attacker.invalid/?token=header-secret",
|
||||
"Bearer body-secret at http://10.0.0.8/private\r\nx-injected: true",
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(TUNNEL_ERROR_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("relay")
|
||||
);
|
||||
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("relay error response body should read");
|
||||
assert_eq!(body.as_ref(), b"tunnel relay failed");
|
||||
let body = String::from_utf8_lossy(&body);
|
||||
assert!(!body.contains("body-secret"));
|
||||
assert!(!body.contains("10.0.0.8"));
|
||||
assert!(!body.contains("x-injected"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rollout_probe_marker_is_only_trusted_from_a_forwarding_gateway() {
|
||||
let mut headers = HeaderMap::new();
|
||||
@@ -522,6 +551,18 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
async fn authenticated_request(envelope: Vec<u8>) -> Request {
|
||||
let spool = crate::tunnel::prepare_owner_relay_request_body(Body::from(envelope))
|
||||
.await
|
||||
.expect("relay envelope should prepare");
|
||||
let mut request = Request::builder()
|
||||
.body(Body::empty())
|
||||
.expect("request should build");
|
||||
request.extensions_mut().insert(RelayRequestAuthenticated);
|
||||
request.extensions_mut().insert(spool);
|
||||
request
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
||||
let meta = protocol::RequestMeta {
|
||||
@@ -591,7 +632,12 @@ mod tests {
|
||||
None,
|
||||
Some(1_800_000_000),
|
||||
None,
|
||||
None,
|
||||
Some(json!({
|
||||
"tunnel_security": {
|
||||
"mode": TUNNEL_SECURITY_NON_TLS_REQUIRED,
|
||||
"encryption_key": LOCAL_TUNNEL_TEST_PSK,
|
||||
}
|
||||
})),
|
||||
None,
|
||||
None,
|
||||
Some(1_800_000_000),
|
||||
@@ -599,6 +645,17 @@ mod tests {
|
||||
Some(1_800_000_000),
|
||||
Some(1_800_000_000),
|
||||
)
|
||||
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||
}
|
||||
|
||||
async fn recv_tunnel_test_frame(
|
||||
proxy_rx: &mut aether_runtime::BoundedQueueReceiver<Message>,
|
||||
description: &str,
|
||||
) -> Message {
|
||||
tokio::time::timeout(Duration::from_secs(5), proxy_rx.recv())
|
||||
.await
|
||||
.unwrap_or_else(|_| panic!("timed out waiting for {description}"))
|
||||
.unwrap_or_else(|| panic!("proxy channel closed before {description}"))
|
||||
}
|
||||
|
||||
fn encode_relay_envelope(meta: &protocol::RequestMeta, body: &[u8]) -> Vec<u8> {
|
||||
@@ -611,9 +668,26 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_rejects_non_loopback_without_forwarded_header() {
|
||||
async fn relay_rejects_unsigned_request_even_from_loopback() {
|
||||
let request = Request::builder()
|
||||
.body(Body::empty())
|
||||
.body(Body::from(encode_relay_envelope(
|
||||
&protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "GET".to_string(),
|
||||
url: "https://example.com/".to_string(),
|
||||
headers: HashMap::new(),
|
||||
stream: false,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
},
|
||||
&[],
|
||||
)))
|
||||
.expect("request should build");
|
||||
let response = relay_request(
|
||||
Path("node-123".to_string()),
|
||||
@@ -625,13 +699,40 @@ mod tests {
|
||||
.into_response();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(TUNNEL_ERROR_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("forbidden")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_accepts_forwarded_gateway_request_from_non_loopback() {
|
||||
async fn relay_rejects_forged_forwarded_gateway_header() {
|
||||
let request = Request::builder()
|
||||
.header(TUNNEL_RELAY_FORWARDED_BY_HEADER, "gateway-a")
|
||||
.body(Body::empty())
|
||||
.header(
|
||||
aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER,
|
||||
"gateway-a",
|
||||
)
|
||||
.body(Body::from(encode_relay_envelope(
|
||||
&protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
method: "GET".to_string(),
|
||||
url: "https://example.com/".to_string(),
|
||||
headers: HashMap::new(),
|
||||
stream: false,
|
||||
request_timeout_ms: None,
|
||||
stream_first_byte_timeout_ms: None,
|
||||
timeout: 30,
|
||||
follow_redirects: None,
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
},
|
||||
&[],
|
||||
)))
|
||||
.expect("request should build");
|
||||
let response = relay_request(
|
||||
Path("node-123".to_string()),
|
||||
@@ -642,24 +743,29 @@ mod tests {
|
||||
.await
|
||||
.into_response();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(TUNNEL_ERROR_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("bad_request")
|
||||
Some("forbidden")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_records_real_traffic_confirmation_for_upgrade_rollout() {
|
||||
let mut node = sample_connected_proxy_node("node-123");
|
||||
node.proxy_metadata = Some(json!({"version": "1.0.0"}));
|
||||
node.proxy_metadata
|
||||
.as_mut()
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
.expect("proxy metadata should be an object")
|
||||
.insert("version".to_string(), json!("1.0.0"));
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![node]));
|
||||
let data = Arc::new(
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(Arc::clone(&repository))
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new()),
|
||||
.with_system_config_values_for_tests(Vec::<(String, serde_json::Value)>::new())
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
|
||||
let started = start_proxy_upgrade_rollout(data.as_ref(), "2.0.0".to_string(), 1, 0, None)
|
||||
@@ -670,6 +776,7 @@ mod tests {
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-123".to_string(),
|
||||
expected_tunnel_generation: None,
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(1),
|
||||
total_requests_delta: Some(1),
|
||||
@@ -677,7 +784,13 @@ mod tests {
|
||||
failed_requests_delta: Some(0),
|
||||
dns_failures_delta: Some(0),
|
||||
stream_errors_delta: Some(0),
|
||||
proxy_metadata: Some(json!({"version": "2.0.0"})),
|
||||
proxy_metadata: Some(json!({
|
||||
"version": "2.0.0",
|
||||
"tunnel_security": {
|
||||
"mode": TUNNEL_SECURITY_NON_TLS_REQUIRED,
|
||||
"encryption_key": LOCAL_TUNNEL_TEST_PSK,
|
||||
}
|
||||
})),
|
||||
proxy_version: Some("2.0.0".to_string()),
|
||||
})
|
||||
.await
|
||||
@@ -692,15 +805,19 @@ mod tests {
|
||||
let state = test_app_state().with_data(Arc::clone(&data));
|
||||
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
state.hub.register_proxy(Arc::new(ProxyConn::new(
|
||||
500,
|
||||
"node-123".to_string(),
|
||||
"Node 123".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
2,
|
||||
)));
|
||||
state.hub.register_proxy(Arc::new(
|
||||
ProxyConn::new(
|
||||
500,
|
||||
"node-123".to_string(),
|
||||
"Node 123".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
2,
|
||||
)
|
||||
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()),
|
||||
));
|
||||
|
||||
let meta = protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
@@ -717,9 +834,7 @@ mod tests {
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
};
|
||||
let request = Request::builder()
|
||||
.body(Body::from(encode_relay_envelope(&meta, &[])))
|
||||
.expect("request should build");
|
||||
let request = authenticated_request(encode_relay_envelope(&meta, &[])).await;
|
||||
|
||||
let relay_state = state.clone();
|
||||
let relay_task = tokio::spawn(async move {
|
||||
@@ -733,7 +848,7 @@ mod tests {
|
||||
.into_response()
|
||||
});
|
||||
|
||||
let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") {
|
||||
let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await {
|
||||
Message::Binary(data) => data,
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
@@ -741,7 +856,7 @@ mod tests {
|
||||
.expect("request header frame should parse");
|
||||
assert_eq!(request_header.msg_type, protocol::REQUEST_HEADERS);
|
||||
|
||||
let request_body = match proxy_rx.recv().await.expect("body frame should arrive") {
|
||||
let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await {
|
||||
Message::Binary(data) => data,
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
@@ -807,18 +922,29 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn relay_strips_hop_by_hop_and_stale_length_headers_from_proxy_response() {
|
||||
let state = test_app_state();
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
sample_connected_proxy_node("node-123"),
|
||||
]));
|
||||
let data = Arc::new(
|
||||
GatewayDataState::with_proxy_node_repository_for_tests(repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let state = test_app_state().with_data(data);
|
||||
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
state.hub.register_proxy(Arc::new(ProxyConn::new(
|
||||
501,
|
||||
"node-123".to_string(),
|
||||
"Node 123".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
2,
|
||||
)));
|
||||
state.hub.register_proxy(Arc::new(
|
||||
ProxyConn::new(
|
||||
501,
|
||||
"node-123".to_string(),
|
||||
"Node 123".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
2,
|
||||
)
|
||||
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string()),
|
||||
));
|
||||
|
||||
let meta = protocol::RequestMeta {
|
||||
provider_id: None,
|
||||
@@ -835,9 +961,7 @@ mod tests {
|
||||
http1_only: false,
|
||||
transport_profile: None,
|
||||
};
|
||||
let request = Request::builder()
|
||||
.body(Body::from(encode_relay_envelope(&meta, &[])))
|
||||
.expect("request should build");
|
||||
let request = authenticated_request(encode_relay_envelope(&meta, &[])).await;
|
||||
|
||||
let relay_state = state.clone();
|
||||
let relay_task = tokio::spawn(async move {
|
||||
@@ -851,7 +975,7 @@ mod tests {
|
||||
.into_response()
|
||||
});
|
||||
|
||||
let request_headers = match proxy_rx.recv().await.expect("headers frame should arrive") {
|
||||
let request_headers = match recv_tunnel_test_frame(&mut proxy_rx, "headers frame").await {
|
||||
Message::Binary(data) => data,
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
@@ -859,7 +983,7 @@ mod tests {
|
||||
.expect("request header frame should parse");
|
||||
assert_eq!(request_header.msg_type, protocol::REQUEST_HEADERS);
|
||||
|
||||
let request_body = match proxy_rx.recv().await.expect("body frame should arrive") {
|
||||
let request_body = match recv_tunnel_test_frame(&mut proxy_rx, "body frame").await {
|
||||
Message::Binary(data) => data,
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
@@ -873,6 +997,17 @@ mod tests {
|
||||
("content-length".to_string(), "999".to_string()),
|
||||
("transfer-encoding".to_string(), "chunked".to_string()),
|
||||
("connection".to_string(), "keep-alive".to_string()),
|
||||
(
|
||||
"connection".to_string(),
|
||||
"x-hop-private, x-accel-redirect".to_string(),
|
||||
),
|
||||
("x-hop-private".to_string(), "secret".to_string()),
|
||||
("x-accel-redirect".to_string(), "/internal".to_string()),
|
||||
("set-cookie".to_string(), "session=attacker".to_string()),
|
||||
(
|
||||
"x-aether-future-control".to_string(),
|
||||
"attacker".to_string(),
|
||||
),
|
||||
("content-type".to_string(), "text/plain".to_string()),
|
||||
(
|
||||
"x-proxy-timing".to_string(),
|
||||
@@ -905,6 +1040,10 @@ mod tests {
|
||||
assert!(response.headers().get("content-length").is_none());
|
||||
assert!(response.headers().get("transfer-encoding").is_none());
|
||||
assert!(response.headers().get("connection").is_none());
|
||||
assert!(response.headers().get("x-hop-private").is_none());
|
||||
assert!(response.headers().get("x-accel-redirect").is_none());
|
||||
assert!(response.headers().get("set-cookie").is_none());
|
||||
assert!(response.headers().get("x-aether-future-control").is_none());
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -11,44 +11,80 @@ use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::watch;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
|
||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
||||
use super::protocol;
|
||||
use aether_contracts::tunnel::Frame;
|
||||
use aether_contracts::tunnel::{Frame, HelloPayload, MsgType};
|
||||
use aether_contracts::tunnel_security::{SecureFrameCodec, TunnelSecurityRole};
|
||||
|
||||
/// Maximum single frame size: 64 MB
|
||||
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
|
||||
/// A connection that has passed the HTTP proof must still complete the
|
||||
/// encrypted protocol handshake promptly. Keeping this deadline independent
|
||||
/// from the normal idle timeout prevents half-open authenticated sockets from
|
||||
/// holding an admission permit indefinitely.
|
||||
const PROXY_HELLO_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
const MAX_PREAUTH_PINGS: usize = 8;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ProxyHelloValidationError {
|
||||
MalformedFrame,
|
||||
DecryptionFailed,
|
||||
UnexpectedFrame,
|
||||
InvalidPayload,
|
||||
ProtocolVersionMismatch,
|
||||
SecuritySessionMismatch,
|
||||
}
|
||||
|
||||
pub async fn handle_proxy_connection(
|
||||
ws: WebSocket,
|
||||
hub: Arc<HubRouter>,
|
||||
node_id: String,
|
||||
node_name: String,
|
||||
node_generation: String,
|
||||
max_streams: usize,
|
||||
protocol_version: u8,
|
||||
security_key: Option<String>,
|
||||
security_session: String,
|
||||
management_token_credential: Option<ProxyManagementTokenCredential>,
|
||||
cfg: ConnConfig,
|
||||
) {
|
||||
let conn_id = hub.alloc_conn_id();
|
||||
let (mut ws_tx, ws_rx) = ws.split();
|
||||
let (mut ws_tx, mut 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 security = match security_key.as_deref() {
|
||||
Some(key) => {
|
||||
match SecureFrameCodec::new(key, &security_session, TunnelSecurityRole::Server) {
|
||||
Ok(codec) => Some(Arc::new(codec)),
|
||||
let (security, initial_hello) = match security_key.as_deref() {
|
||||
Some(security_key) => {
|
||||
let security = match SecureFrameCodec::new(
|
||||
security_key,
|
||||
&security_session,
|
||||
TunnelSecurityRole::Server,
|
||||
) {
|
||||
Ok(codec) => Arc::new(codec),
|
||||
Err(error) => {
|
||||
warn!(conn_id, node_id = %node_id, error = %error, "secure tunnel codec initialization failed");
|
||||
return;
|
||||
}
|
||||
}
|
||||
};
|
||||
let Some(hello) = read_authenticated_proxy_hello(
|
||||
&mut ws_tx,
|
||||
&mut ws_rx,
|
||||
security.as_ref(),
|
||||
protocol_version,
|
||||
&security_session,
|
||||
conn_id,
|
||||
&node_id,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return;
|
||||
};
|
||||
(Some(security), Some(hello))
|
||||
}
|
||||
None => None,
|
||||
None => (None, None),
|
||||
};
|
||||
|
||||
let conn = Arc::new(ProxyConn::new(
|
||||
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
let conn = ProxyConn::new(
|
||||
conn_id,
|
||||
node_id.clone(),
|
||||
node_name.clone(),
|
||||
@@ -56,9 +92,21 @@ pub async fn handle_proxy_connection(
|
||||
close_tx,
|
||||
max_streams,
|
||||
protocol_version,
|
||||
));
|
||||
)
|
||||
.with_tunnel_generation(node_generation);
|
||||
let conn = match (security_key.clone(), management_token_credential) {
|
||||
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
||||
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
||||
(Some(_), Some(_)) | (None, None) => {
|
||||
warn!(conn_id, node_id = %node_id, "proxy connection missing an unambiguous credential binding");
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
hub.register_proxy(conn.clone());
|
||||
if let Some(mut hello) = initial_hello {
|
||||
hub.handle_proxy_frame(conn.id, &mut hello).await;
|
||||
}
|
||||
|
||||
let writer_conn_id = conn_id;
|
||||
let writer_conn = conn.clone();
|
||||
@@ -75,7 +123,7 @@ pub async fn handle_proxy_connection(
|
||||
_ => 0,
|
||||
};
|
||||
let send_started_at = std::time::Instant::now();
|
||||
let msg = match encrypt_message(msg, writer_security.as_deref()) {
|
||||
let msg = match encrypt_message(msg, writer_security.as_deref()) {
|
||||
Ok(msg) => msg,
|
||||
Err(error) => {
|
||||
warn!(conn_id = writer_conn_id, error = %error, "failed to encrypt outbound proxy frame");
|
||||
@@ -244,6 +292,120 @@ pub async fn handle_proxy_connection(
|
||||
let _ = writer.await;
|
||||
}
|
||||
|
||||
async fn read_authenticated_proxy_hello(
|
||||
ws_tx: &mut futures_util::stream::SplitSink<WebSocket, Message>,
|
||||
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
|
||||
security: &SecureFrameCodec,
|
||||
protocol_version: u8,
|
||||
security_session: &str,
|
||||
conn_id: u64,
|
||||
node_id: &str,
|
||||
) -> Option<Vec<u8>> {
|
||||
let result = tokio::time::timeout(PROXY_HELLO_TIMEOUT, async {
|
||||
let mut preauth_pings = 0usize;
|
||||
loop {
|
||||
match ws_rx.next().await {
|
||||
Some(Ok(Message::Binary(data))) => {
|
||||
return match validate_authenticated_proxy_hello(
|
||||
data,
|
||||
security,
|
||||
protocol_version,
|
||||
security_session,
|
||||
) {
|
||||
Ok(hello) => Some(hello),
|
||||
Err(error) => {
|
||||
warn!(
|
||||
conn_id,
|
||||
node_id = %node_id,
|
||||
?error,
|
||||
"proxy connection rejected: invalid encrypted HELLO"
|
||||
);
|
||||
None
|
||||
}
|
||||
};
|
||||
}
|
||||
Some(Ok(Message::Ping(payload))) => {
|
||||
preauth_pings = preauth_pings.saturating_add(1);
|
||||
if preauth_pings > MAX_PREAUTH_PINGS {
|
||||
warn!(
|
||||
conn_id,
|
||||
node_id = %node_id,
|
||||
"proxy connection rejected: too many WebSocket pings before encrypted HELLO"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
if let Err(error) = ws_tx.send(Message::Pong(payload)).await {
|
||||
warn!(conn_id, node_id = %node_id, error = %error, "failed to answer WebSocket ping before proxy authentication");
|
||||
return None;
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Pong(_))) => {}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(conn_id, node_id = %node_id, "proxy disconnected before encrypted HELLO authentication");
|
||||
return None;
|
||||
}
|
||||
Some(Ok(Message::Text(_))) => {
|
||||
warn!(conn_id, node_id = %node_id, "proxy connection rejected: text message received before encrypted HELLO");
|
||||
return None;
|
||||
}
|
||||
Some(Err(error)) => {
|
||||
warn!(conn_id, node_id = %node_id, error = %error, "proxy WebSocket failed before encrypted HELLO authentication");
|
||||
return None;
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(hello) => hello,
|
||||
Err(_) => {
|
||||
warn!(
|
||||
conn_id,
|
||||
node_id = %node_id,
|
||||
timeout_ms = PROXY_HELLO_TIMEOUT.as_millis(),
|
||||
"proxy connection rejected: encrypted HELLO timed out"
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_authenticated_proxy_hello(
|
||||
data: bytes::Bytes,
|
||||
security: &SecureFrameCodec,
|
||||
protocol_version: u8,
|
||||
security_session: &str,
|
||||
) -> Result<Vec<u8>, ProxyHelloValidationError> {
|
||||
let header =
|
||||
protocol::FrameHeader::parse(&data).ok_or(ProxyHelloValidationError::MalformedFrame)?;
|
||||
let expected_len = protocol::HEADER_SIZE
|
||||
.checked_add(header.payload_len as usize)
|
||||
.ok_or(ProxyHelloValidationError::MalformedFrame)?;
|
||||
if expected_len != data.len() {
|
||||
return Err(ProxyHelloValidationError::MalformedFrame);
|
||||
}
|
||||
|
||||
let frame = Frame::decode(data).map_err(|_| ProxyHelloValidationError::MalformedFrame)?;
|
||||
let frame = security
|
||||
.decrypt_frame(frame)
|
||||
.map_err(|_| ProxyHelloValidationError::DecryptionFailed)?;
|
||||
if frame.stream_id != 0 || frame.msg_type != MsgType::Hello || frame.flags != 0 {
|
||||
return Err(ProxyHelloValidationError::UnexpectedFrame);
|
||||
}
|
||||
|
||||
let hello = serde_json::from_slice::<HelloPayload>(&frame.payload)
|
||||
.map_err(|_| ProxyHelloValidationError::InvalidPayload)?;
|
||||
if hello.protocol_version != protocol_version {
|
||||
return Err(ProxyHelloValidationError::ProtocolVersionMismatch);
|
||||
}
|
||||
if hello.session_id.as_deref() != Some(security_session) {
|
||||
return Err(ProxyHelloValidationError::SecuritySessionMismatch);
|
||||
}
|
||||
|
||||
Ok(frame.encode().to_vec())
|
||||
}
|
||||
|
||||
async fn run_proxy_reader(
|
||||
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
|
||||
hub: Arc<HubRouter>,
|
||||
@@ -358,3 +520,152 @@ fn decrypt_message(
|
||||
let frame = codec.decrypt_frame(frame)?;
|
||||
Ok(frame.encode().to_vec())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||
const SESSION: &str = "0123456789abcdef0123456789abcdef";
|
||||
const PROTOCOL_VERSION: u8 = 3;
|
||||
|
||||
fn codecs() -> (SecureFrameCodec, SecureFrameCodec) {
|
||||
(
|
||||
SecureFrameCodec::new(KEY, SESSION, TunnelSecurityRole::Client).expect("client codec"),
|
||||
SecureFrameCodec::new(KEY, SESSION, TunnelSecurityRole::Server).expect("server codec"),
|
||||
)
|
||||
}
|
||||
|
||||
fn hello_frame(protocol_version: u8, session_id: &str) -> Frame {
|
||||
Frame::control(
|
||||
MsgType::Hello,
|
||||
serde_json::to_vec(&HelloPayload {
|
||||
protocol_version,
|
||||
capabilities: vec!["flow-control".to_string()],
|
||||
session_id: Some(session_id.to_string()),
|
||||
replica_id: None,
|
||||
})
|
||||
.expect("hello payload"),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_proxy_hello_accepts_bound_encrypted_control_frame() {
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION))
|
||||
.expect("encrypted HELLO");
|
||||
|
||||
let clear =
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION)
|
||||
.expect("authenticated HELLO");
|
||||
let frame = Frame::decode(bytes::Bytes::from(clear)).expect("clear HELLO");
|
||||
|
||||
assert_eq!(frame.stream_id, 0);
|
||||
assert_eq!(frame.msg_type, MsgType::Hello);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_proxy_hello_advances_shared_receive_sequence() {
|
||||
let (client, server) = codecs();
|
||||
let encrypted_hello = client
|
||||
.encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION))
|
||||
.expect("encrypted HELLO");
|
||||
validate_authenticated_proxy_hello(encrypted_hello, &server, PROTOCOL_VERSION, SESSION)
|
||||
.expect("authenticated HELLO");
|
||||
|
||||
let encrypted_settings = client
|
||||
.encrypt_frame(Frame::control(MsgType::Settings, bytes::Bytes::new()))
|
||||
.expect("encrypted SETTINGS");
|
||||
let settings = server
|
||||
.decrypt_frame(Frame::decode(encrypted_settings).expect("wire SETTINGS"))
|
||||
.expect("next sequence should decrypt");
|
||||
|
||||
assert_eq!(settings.msg_type, MsgType::Settings);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_proxy_hello_rejects_non_encrypted_or_wrong_frame() {
|
||||
let (_, server) = codecs();
|
||||
let clear = hello_frame(PROTOCOL_VERSION, SESSION).encode();
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(clear, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::DecryptionFailed)
|
||||
);
|
||||
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(Frame::control(MsgType::Settings, bytes::Bytes::new()))
|
||||
.expect("encrypted SETTINGS");
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::UnexpectedFrame)
|
||||
);
|
||||
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(Frame::new(
|
||||
1,
|
||||
MsgType::Hello,
|
||||
0,
|
||||
hello_frame(PROTOCOL_VERSION, SESSION).payload,
|
||||
))
|
||||
.expect("encrypted stream HELLO");
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::UnexpectedFrame)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_proxy_hello_rejects_protocol_or_session_mismatch() {
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(hello_frame(PROTOCOL_VERSION - 1, SESSION))
|
||||
.expect("encrypted HELLO");
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::ProtocolVersionMismatch)
|
||||
);
|
||||
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(hello_frame(PROTOCOL_VERSION, "different-session"))
|
||||
.expect("encrypted HELLO");
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::SecuritySessionMismatch)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_proxy_hello_rejects_malformed_or_ambiguous_frame() {
|
||||
let (client, server) = codecs();
|
||||
let mut encrypted = client
|
||||
.encrypt_frame(hello_frame(PROTOCOL_VERSION, SESSION))
|
||||
.expect("encrypted HELLO")
|
||||
.to_vec();
|
||||
encrypted.push(0);
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(
|
||||
bytes::Bytes::from(encrypted),
|
||||
&server,
|
||||
PROTOCOL_VERSION,
|
||||
SESSION,
|
||||
),
|
||||
Err(ProxyHelloValidationError::MalformedFrame)
|
||||
);
|
||||
|
||||
let (client, server) = codecs();
|
||||
let encrypted = client
|
||||
.encrypt_frame(Frame::control(
|
||||
MsgType::Hello,
|
||||
bytes::Bytes::from_static(b"not-json"),
|
||||
))
|
||||
.expect("encrypted malformed HELLO");
|
||||
assert_eq!(
|
||||
validate_authenticated_proxy_hello(encrypted, &server, PROTOCOL_VERSION, SESSION),
|
||||
Err(ProxyHelloValidationError::InvalidPayload)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+2336
-207
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user