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:
elky
2026-09-04 03:45:52 +08:00
parent ddcbeb3ae9
commit 579f2c7cc1
1019 changed files with 190437 additions and 26080 deletions
+36 -11
View File
@@ -19,8 +19,8 @@ use tokio::task::JoinHandle;
use tracing::{error, info, warn};
use crate::config::{
effective_tunnel_security, validate_tunnel_encryption_key, Config, ServerEntry,
TunnelPoolSizing,
aether_url_for_log, effective_tunnel_security, validate_tunnel_encryption_key, Config,
ServerEntry, TunnelPoolSizing,
};
use crate::net;
use crate::registration::client::AetherClient;
@@ -96,6 +96,11 @@ struct DiagnosticsState {
/// Run the full application lifecycle after config has been parsed.
pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Result<()> {
config.validate()?;
for (index, server) in servers.iter().enumerate() {
server
.validate()
.map_err(|error| anyhow::anyhow!("servers[{index}] invalid: {error}"))?;
}
init_tracing(&config);
info!(
@@ -227,7 +232,7 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
if entry.aether_url.trim_start().starts_with("http://") {
warn!(
server = %label,
url = %entry.aether_url,
url = %aether_url_for_log(&entry.aether_url),
"secure tunnel frame encryption starts after registration; deliver install and registration credentials over HTTPS or another trusted bootstrap channel"
);
}
@@ -245,16 +250,29 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
.register(&config, entry, &node_name, &public_ip, Some(&hw_info))
.await
{
Ok(node_id) => {
info!(server = %label, node_id = %node_id, url = %entry.aether_url, node_name = %node_name, "registered");
Ok(registration) => {
info!(
server = %label,
node_id = %registration.node_id,
tunnel_generation = %registration.tunnel_generation,
url = %aether_url_for_log(&entry.aether_url),
node_name = %node_name,
"registered"
);
server_contexts.lock().await.push(build_server_context(
&config, &label, entry, client, &node_name, node_id,
&config,
&label,
entry,
client,
&node_name,
registration.node_id,
registration.tunnel_generation,
));
}
Err(e) => {
warn!(
server = %label,
url = %entry.aether_url,
url = %aether_url_for_log(&entry.aether_url),
error = %e,
"registration failed, will retry in background"
);
@@ -644,15 +662,16 @@ async fn retry_failed_registration(
)
.await
{
Ok(node_id) => {
info!(server = %label, node_id = %node_id, attempt, "registration retry succeeded");
Ok(registration) => {
info!(server = %label, node_id = %registration.node_id, tunnel_generation = %registration.tunnel_generation, attempt, "registration retry succeeded");
let server = build_server_context(
&state.config,
&label,
&entry,
client,
&node_name,
node_id,
registration.node_id,
registration.tunnel_generation,
);
server_contexts.lock().await.push(Arc::clone(&server));
spawn_tunnel_pool_manager(
@@ -712,6 +731,7 @@ fn build_server_context(
client: Arc<AetherClient>,
node_name: &str,
node_id: String,
tunnel_generation: String,
) -> Arc<ServerContext> {
let mut dynamic = DynamicConfig::from_config(config);
dynamic.node_name = node_name.to_string();
@@ -727,6 +747,7 @@ fn build_server_context(
tunnel_encryption_key: entry.tunnel_encryption_key.clone(),
node_name: node_name.to_string(),
node_id: Arc::new(RwLock::new(node_id)),
tunnel_generation,
aether_client: client,
dynamic: Arc::new(ArcSwap::from_pointee(dynamic)),
active_connections: Arc::new(AtomicU64::new(0)),
@@ -1290,7 +1311,10 @@ mod tests {
register_hits.fetch_add(1, Ordering::SeqCst);
(
AxumStatusCode::OK,
axum::Json(json!({ "node_id": "node-recovery" })),
axum::Json(json!({
"node_id": "node-recovery",
"tunnel_generation": "test-generation-recovery"
})),
)
}
@@ -1351,6 +1375,7 @@ mod tests {
client,
&state.config.node_name,
node_id.to_string(),
"test-generation-1".to_string(),
)
}
+595 -13
View File
@@ -1,4 +1,5 @@
use std::fmt;
use std::io::{self, Read, Write};
use std::net::SocketAddr;
use std::path::Path;
use std::str::FromStr;
@@ -79,6 +80,12 @@ const TUNNEL_PROFILE_ENV: &str = "AETHER_TUNNEL_PROFILE";
const TUNNEL_STREAM_INITIAL_WINDOW_BYTES_ENV: &str = "AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES";
const TUNNEL_DRAIN_DEADLINE_MS_ENV: &str = "AETHER_TUNNEL_DRAIN_DEADLINE_MS";
// The configuration contains only scalar settings and a bounded list of
// server entries. Refuse an unexpectedly large local file before TOML parsing
// so a replaced or corrupted config cannot force an unbounded allocation at
// service startup.
const MAX_CONFIG_FILE_BYTES: u64 = 1024 * 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TunnelPoolSizing {
pub initial_connections: u32,
@@ -193,12 +200,43 @@ pub fn effective_tunnel_security(
TunnelSecurity::Off
}
pub(crate) fn validate_aether_url(value: &str) -> anyhow::Result<()> {
let value = value.trim();
let parsed = url::Url::parse(value)
.map_err(|_| anyhow::anyhow!("aether_url must be an absolute HTTP(S) URL"))?;
if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() {
anyhow::bail!("aether_url must be an absolute HTTP(S) URL");
}
if !aether_http::is_https_or_loopback_http_url(&parsed) {
anyhow::bail!(
"aether_url must use HTTPS; HTTP is allowed only for a literal loopback host"
);
}
if !parsed.username().is_empty() || parsed.password().is_some() {
anyhow::bail!("aether_url must not contain embedded credentials");
}
if parsed.query().is_some() || parsed.fragment().is_some() {
anyhow::bail!("aether_url must not contain a query string or fragment");
}
Ok(())
}
pub(crate) fn aether_url_for_log(value: &str) -> String {
let Ok(parsed) = url::Url::parse(value.trim()) else {
return "<invalid-aether-url>".to_string();
};
if !matches!(parsed.scheme(), "http" | "https" | "ws" | "wss") || parsed.host_str().is_none() {
return "<invalid-aether-url>".to_string();
}
parsed.origin().ascii_serialization()
}
/// Aether tunnel agent.
///
/// Deployed on overseas VPS to relay API traffic for Aether instances
/// behind the GFW. Connects to Aether via WebSocket tunnel, registers
/// with Aether, and relays upstream requests.
#[derive(Parser, Debug, Clone)]
#[derive(Parser, Clone)]
#[command(version, about)]
pub struct Config {
/// Aether server URL (e.g. https://aether.example.com)
@@ -250,11 +288,12 @@ pub struct Config {
)]
pub allowed_ports: Vec<u16>,
/// Allow private/reserved upstream IP targets. Enabled by default.
/// Allow private/reserved upstream IP targets. Disabled by default; enable
/// explicitly only for deployments that require access to private services.
#[arg(
long,
env = "AETHER_TUNNEL_ALLOW_PRIVATE_TARGETS",
default_value_t = true
default_value_t = false
)]
pub allow_private_targets: bool,
@@ -644,10 +683,31 @@ pub struct Config {
pub tunnel_scale_down_grace_secs: u64,
}
impl std::fmt::Debug for Config {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("Config")
.field("aether_url", &aether_url_for_log(&self.aether_url))
.field("management_token", &"<redacted>")
.field("node_name", &self.node_name)
.field("node_region", &self.node_region)
.field("tunnel_security", &self.tunnel_security)
.field(
"tunnel_encryption_key",
&self.tunnel_encryption_key.as_ref().map(|_| "<redacted>"),
)
.finish_non_exhaustive()
}
}
impl Config {
/// Validate configuration values are within sane ranges.
/// Called after parsing to catch misconfigurations early.
pub fn validate(&self) -> anyhow::Result<()> {
validate_aether_url(&self.aether_url)?;
if self.management_token.trim().is_empty() {
anyhow::bail!("management_token must not be empty");
}
if self.heartbeat_interval == 0 {
anyhow::bail!("heartbeat_interval must be > 0");
}
@@ -698,6 +758,9 @@ impl Config {
if matches!(self.tunnel_connections_max, Some(0)) {
anyhow::bail!("tunnel_connections_max must be > 0");
}
if matches!(self.tunnel_max_streams, Some(0)) {
anyhow::bail!("tunnel_max_streams must be > 0");
}
if self.tunnel_stream_initial_window_bytes == 0 {
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
}
@@ -906,7 +969,7 @@ impl Config {
}
/// Per-server connection config (used in multi-server TOML `[[servers]]`).
#[derive(Debug, Clone, Serialize, Deserialize)]
#[derive(Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ServerEntry {
pub aether_url: String,
@@ -921,13 +984,39 @@ pub struct ServerEntry {
pub tunnel_encryption_key: Option<String>,
}
impl std::fmt::Debug for ServerEntry {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ServerEntry")
.field("aether_url", &aether_url_for_log(&self.aether_url))
.field("management_token", &"<redacted>")
.field("node_name", &self.node_name)
.field("tunnel_security", &self.tunnel_security)
.field(
"tunnel_encryption_key",
&self.tunnel_encryption_key.as_ref().map(|_| "<redacted>"),
)
.finish()
}
}
impl ServerEntry {
pub(crate) fn validate(&self) -> anyhow::Result<()> {
validate_aether_url(&self.aether_url)?;
if self.management_token.trim().is_empty() {
anyhow::bail!("management_token must not be empty");
}
Ok(())
}
}
// ---------------------------------------------------------------------------
// TOML config file support
// ---------------------------------------------------------------------------
/// Serializable config for TOML file persistence.
/// All fields are optional -- only populated values are written.
#[derive(Debug, Default, Serialize, Deserialize)]
#[derive(Default, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ConfigFile {
#[serde(skip_serializing_if = "Option::is_none")]
@@ -1048,17 +1137,46 @@ pub struct ConfigFile {
pub servers: Vec<ServerEntry>,
}
impl std::fmt::Debug for ConfigFile {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ConfigFile")
.field("node_name", &self.node_name)
.field("node_region", &self.node_region)
.field("server_count", &self.servers.len())
.finish_non_exhaustive()
}
}
impl ConfigFile {
/// Load from a TOML file.
pub fn load(path: &Path) -> anyhow::Result<Self> {
let content = std::fs::read_to_string(path)?;
let mut file = std::fs::File::open(path)?;
let advertised_len = file.metadata()?.len();
if advertised_len > MAX_CONFIG_FILE_BYTES {
anyhow::bail!(
"tunnel config exceeds the {} byte limit",
MAX_CONFIG_FILE_BYTES
);
}
let mut content = String::with_capacity(advertised_len as usize);
Read::by_ref(&mut file)
.take(MAX_CONFIG_FILE_BYTES.saturating_add(1))
.read_to_string(&mut content)?;
if content.len() as u64 > MAX_CONFIG_FILE_BYTES {
anyhow::bail!(
"tunnel config exceeds the {} byte limit",
MAX_CONFIG_FILE_BYTES
);
}
parse_config_file_content(&content)
}
/// Save to a TOML file.
pub fn save(&self, path: &Path) -> anyhow::Result<()> {
let content = toml::to_string_pretty(self)?;
std::fs::write(path, content)?;
write_private_config_atomically(path, content.as_bytes())?;
Ok(())
}
@@ -1268,6 +1386,155 @@ impl ConfigFile {
}
}
fn write_private_config_atomically(path: &Path, content: &[u8]) -> io::Result<()> {
reject_config_symlink(path)?;
let parent = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."));
let file_name = path.file_name().ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
"configuration path must name a file",
)
})?;
if !parent.is_dir() {
return Err(io::Error::new(
io::ErrorKind::NotFound,
"configuration parent directory does not exist",
));
}
let mut temp_path = None;
let mut temp_file = None;
for _ in 0..16 {
let candidate = parent.join(format!(
".{}.tmp-{}",
file_name.to_string_lossy(),
uuid::Uuid::new_v4()
));
let mut options = std::fs::OpenOptions::new();
options.write(true).create_new(true);
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt as _;
options.mode(0o600);
}
match options.open(&candidate) {
Ok(file) => {
temp_path = Some(candidate);
temp_file = Some(file);
break;
}
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue,
Err(error) => return Err(error),
}
}
let temp_path = temp_path.ok_or_else(|| {
io::Error::new(
io::ErrorKind::AlreadyExists,
"could not allocate a unique configuration temporary file",
)
})?;
let mut temp_file = temp_file.expect("temporary path and file are created together");
let replace_result = (|| -> io::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt as _;
temp_file.set_permissions(std::fs::Permissions::from_mode(0o600))?;
}
temp_file.write_all(content)?;
temp_file.sync_all()?;
drop(temp_file);
// Do not silently replace a credential-bearing symlink. The second
// check also covers a target created while the temporary file was written.
reject_config_symlink(path)?;
replace_config_file(&temp_path, path)?;
#[cfg(unix)]
std::fs::File::open(parent)?.sync_all()?;
Ok(())
})();
if replace_result.is_err() {
let _ = std::fs::remove_file(&temp_path);
}
replace_result
}
fn reject_config_symlink(path: &Path) -> io::Result<()> {
match std::fs::symlink_metadata(path) {
Ok(metadata) if config_metadata_is_link_like(&metadata) => Err(io::Error::new(
io::ErrorKind::PermissionDenied,
"refusing to save configuration through a symbolic link or reparse point",
)),
Ok(_) => Ok(()),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error),
}
}
fn config_metadata_is_link_like(metadata: &std::fs::Metadata) -> bool {
#[cfg(windows)]
{
use std::os::windows::fs::MetadataExt as _;
const FILE_ATTRIBUTE_REPARSE_POINT: u32 = 0x0400;
metadata.file_attributes() & FILE_ATTRIBUTE_REPARSE_POINT != 0
}
#[cfg(not(windows))]
{
metadata.file_type().is_symlink()
}
}
#[cfg(not(windows))]
fn replace_config_file(temp_path: &Path, path: &Path) -> io::Result<()> {
std::fs::rename(temp_path, path)
}
#[cfg(windows)]
fn replace_config_file(temp_path: &Path, path: &Path) -> io::Result<()> {
use std::os::windows::ffi::OsStrExt as _;
const MOVEFILE_REPLACE_EXISTING: u32 = 0x0000_0001;
const MOVEFILE_WRITE_THROUGH: u32 = 0x0000_0008;
#[link(name = "kernel32")]
unsafe extern "system" {
fn MoveFileExW(
existing_file_name: *const u16,
new_file_name: *const u16,
flags: u32,
) -> i32;
}
let existing_file_name = temp_path
.as_os_str()
.encode_wide()
.chain(std::iter::once(0))
.collect::<Vec<_>>();
let new_file_name = path
.as_os_str()
.encode_wide()
.chain(std::iter::once(0))
.collect::<Vec<_>>();
let replaced = unsafe {
MoveFileExW(
existing_file_name.as_ptr(),
new_file_name.as_ptr(),
MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH,
)
};
if replaced == 0 {
Err(io::Error::last_os_error())
} else {
Ok(())
}
}
fn parse_config_file_content(content: &str) -> anyhow::Result<ConfigFile> {
reject_removed_config_keys(content)?;
let mut value: toml::Value = toml::from_str(content)?;
@@ -1396,6 +1663,151 @@ mod tests {
use super::*;
use crate::hardware::HardwareInfo;
fn config_save_test_dir(label: &str) -> std::path::PathBuf {
let path = std::env::temp_dir().join(format!(
"aether-tunnel-config-{label}-{}",
uuid::Uuid::new_v4()
));
std::fs::create_dir(&path).expect("config save test directory should be created");
path
}
fn secret_bearing_config_file(secret: &str) -> ConfigFile {
ConfigFile {
node_name: Some("secure-save-test".to_string()),
servers: vec![ServerEntry {
aether_url: "https://example.com".to_string(),
management_token: secret.to_string(),
node_name: None,
tunnel_security: None,
tunnel_encryption_key: None,
}],
..ConfigFile::default()
}
}
#[test]
fn config_file_save_replaces_an_existing_config() {
let directory = config_save_test_dir("replace-existing");
let path = directory.join("tunnel.toml");
secret_bearing_config_file("first-management-secret")
.save(&path)
.expect("initial config save should succeed");
secret_bearing_config_file("second-management-secret")
.save(&path)
.expect("replacement config save should succeed");
let saved = std::fs::read_to_string(&path).expect("replacement config should be readable");
assert!(saved.contains("second-management-secret"));
assert!(!saved.contains("first-management-secret"));
assert_eq!(
std::fs::read_dir(&directory)
.expect("test directory should be readable")
.count(),
1,
"replacement save must not leave a temporary file"
);
std::fs::remove_dir_all(directory).expect("config save test directory should be removed");
}
#[cfg(unix)]
#[test]
fn config_file_save_atomically_replaces_with_owner_only_permissions() {
use std::io::Read as _;
use std::os::unix::fs::PermissionsExt as _;
let directory = config_save_test_dir("atomic-private");
let path = directory.join("tunnel.toml");
std::fs::write(&path, "old configuration")
.expect("existing config fixture should be written");
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o644))
.expect("existing config fixture should be world-readable");
let mut old_handle =
std::fs::File::open(&path).expect("existing config handle should open");
secret_bearing_config_file("management-secret")
.save(&path)
.expect("config save should succeed");
let mode = std::fs::metadata(&path)
.expect("saved config metadata should exist")
.permissions()
.mode()
& 0o777;
assert_eq!(mode, 0o600);
let saved = std::fs::read_to_string(&path).expect("saved config should be readable");
assert!(saved.contains("management-secret"));
let mut old = String::new();
old_handle
.read_to_string(&mut old)
.expect("old inode should remain readable through its open handle");
assert_eq!(old, "old configuration");
assert_eq!(
std::fs::read_dir(&directory)
.expect("test directory should be readable")
.count(),
1,
"successful save must not leave a temporary file"
);
std::fs::remove_dir_all(directory).expect("config save test directory should be removed");
}
#[cfg(unix)]
#[test]
fn config_file_save_rejects_symbolic_link_targets() {
use std::os::unix::fs::symlink;
let directory = config_save_test_dir("symlink");
let target = directory.join("target.toml");
let path = directory.join("tunnel.toml");
std::fs::write(&target, "target sentinel")
.expect("symlink target fixture should be written");
symlink(&target, &path).expect("config symlink fixture should be created");
let error = secret_bearing_config_file("management-secret")
.save(&path)
.expect_err("saving through a symlink must fail");
assert!(error.to_string().contains("symbolic link"));
assert_eq!(
std::fs::read_to_string(&target).expect("symlink target should remain readable"),
"target sentinel"
);
assert!(std::fs::symlink_metadata(&path)
.expect("config symlink should still exist")
.file_type()
.is_symlink());
std::fs::remove_dir_all(directory).expect("config save test directory should be removed");
}
#[test]
fn config_file_save_cleans_temporary_file_when_replace_fails() {
let directory = config_save_test_dir("cleanup");
let path = directory.join("destination-is-a-directory");
std::fs::create_dir(&path).expect("destination directory fixture should be created");
secret_bearing_config_file("management-secret")
.save(&path)
.expect_err("replacing a directory with a config file must fail");
let entries = std::fs::read_dir(&directory)
.expect("test directory should be readable")
.map(|entry| {
entry
.expect("test directory entry should be readable")
.path()
})
.collect::<Vec<_>>();
assert_eq!(entries, vec![path]);
std::fs::remove_dir_all(directory).expect("config save test directory should be removed");
}
#[test]
fn config_file_load_ignores_removed_redirect_replay_budget() {
let config = parse_config_file_content("redirect_replay_budget_bytes = \"1K\"")
@@ -1404,6 +1816,19 @@ mod tests {
assert!(!serialized.contains("redirect_replay_budget_bytes"));
}
#[test]
fn config_file_load_rejects_oversized_files_before_parsing() {
let directory = config_save_test_dir("oversized-load");
let path = directory.join("tunnel.toml");
std::fs::write(&path, vec![b'a'; MAX_CONFIG_FILE_BYTES as usize + 1])
.expect("oversized config fixture should be written");
let error = ConfigFile::load(&path).expect_err("oversized config must be rejected");
assert!(error.to_string().contains("exceeds"));
std::fs::remove_dir_all(directory).expect("config test directory should be removed");
}
#[test]
fn cli_accepts_but_hides_legacy_redirect_replay_budget() {
let config = Config::parse_from([
@@ -1464,7 +1889,7 @@ tunnel_ipv6_only = false
let cfg: ConfigFile = toml::from_str(
r#"
[[servers]]
aether_url = "http://aether.example.com"
aether_url = "http://127.0.0.1:8084"
management_token = "ae_test"
node_name = "jp-proxy-01"
tunnel_security = "non_tls_required"
@@ -1664,7 +2089,7 @@ node_name = "tunnel-test"
}
#[test]
fn cli_defaults_private_targets_to_enabled() {
fn cli_defaults_private_targets_to_disabled() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
@@ -1674,6 +2099,21 @@ node_name = "tunnel-test"
"--node-name",
"tunnel-test",
]);
assert!(!config.allow_private_targets);
}
#[test]
fn cli_allows_explicit_private_targets() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
"--allow-private-targets",
]);
assert!(config.allow_private_targets);
}
@@ -1713,12 +2153,134 @@ node_name = "tunnel-test"
assert!(config.tunnel_encryption_key.is_none());
}
#[test]
fn aether_url_validation_rejects_embedded_secrets_and_non_http_schemes() {
for value in [
"https://alice:[email protected]",
"https://example.com?token=secret",
"https://example.com#secret-fragment",
"http://example.com",
"http://10.0.0.1:8084",
"http://[::ffff:127.0.0.1]:8084",
"file:///etc/passwd",
"not-a-url",
] {
assert!(
validate_aether_url(value).is_err(),
"URL should be rejected: {value}"
);
}
validate_aether_url("https://example.com/base/path")
.expect("ordinary HTTPS URL should validate");
validate_aether_url("http://127.0.0.1:8084/base/path")
.expect("literal loopback HTTP should validate");
validate_aether_url("http://[::1]:8084/base/path")
.expect("literal IPv6 loopback HTTP should validate");
}
#[test]
fn aether_url_log_projection_removes_credentials_query_and_fragment() {
let projected = aether_url_for_log(
"https://alice:[email protected]/base?token=query-secret#secret-fragment",
);
assert_eq!(projected, "https://example.com");
for secret in ["alice", "password", "query-secret", "secret-fragment"] {
assert!(!projected.contains(secret));
}
assert_eq!(aether_url_for_log("not-a-url"), "<invalid-aether-url>");
}
#[test]
fn config_and_server_entries_require_management_tokens() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
" ",
"--node-name",
"tunnel-test",
]);
assert!(config.validate().is_err());
let entry = ServerEntry {
aether_url: "https://example.com".to_string(),
management_token: " ".to_string(),
node_name: None,
tunnel_security: None,
tunnel_encryption_key: None,
};
assert!(entry.validate().is_err());
}
#[test]
fn secret_bearing_config_debug_output_is_redacted() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://alice:[email protected]/base?token=query-secret",
"--management-token",
"management-secret",
"--node-name",
"tunnel-test",
"--tunnel-encryption-key",
"psk-secret",
]);
let config_debug = format!("{config:?}");
for secret in [
"alice",
"password",
"query-secret",
"management-secret",
"psk-secret",
] {
assert!(!config_debug.contains(secret));
}
let entry = ServerEntry {
aether_url:
"https://alice:[email protected]/base?token=query-secret#fragment-secret"
.to_string(),
management_token: "management-secret".to_string(),
node_name: Some("edge-1".to_string()),
tunnel_security: Some(TunnelSecurity::NonTlsRequired),
tunnel_encryption_key: Some("psk-secret".to_string()),
};
let entry_debug = format!("{entry:?}");
for secret in [
"alice",
"password",
"query-secret",
"fragment-secret",
"management-secret",
"psk-secret",
] {
assert!(!entry_debug.contains(secret));
}
let file = ConfigFile {
upstream_proxy_url: Some("http://proxy-user:[email protected]".to_string()),
servers: vec![entry],
..ConfigFile::default()
};
let file_debug = format!("{file:?}");
for secret in [
"proxy-user",
"proxy-secret",
"management-secret",
"psk-secret",
] {
assert!(!file_debug.contains(secret));
}
}
#[test]
fn validate_requires_encryption_key_for_non_tls_security() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
@@ -1735,7 +2297,7 @@ node_name = "tunnel-test"
let with_key = Config::parse_from([
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
@@ -1755,7 +2317,7 @@ node_name = "tunnel-test"
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
@@ -1791,7 +2353,7 @@ node_name = "tunnel-test"
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
@@ -1913,6 +2475,26 @@ node_name = "tunnel-test"
assert!(error.to_string().contains("tunnel_ipv4_only"));
}
#[test]
fn validate_rejects_zero_tunnel_max_streams() {
let config = Config::parse_from([
"aether-tunnel",
"--aether-url",
"https://example.com",
"--management-token",
"ae_test",
"--node-name",
"tunnel-test",
"--tunnel-max-streams",
"0",
]);
let error = config
.validate()
.expect_err("zero tunnel stream capacity must be rejected");
assert!(error.to_string().contains("tunnel_max_streams"));
}
#[test]
fn tunnel_fast_recovery_defaults_use_millisecond_values() {
let config = Config::parse_from([
+107 -15
View File
@@ -40,7 +40,7 @@ pub(crate) enum UpstreamProxyScheme {
Socks5h,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
pub(crate) struct UpstreamProxyConfig {
raw: String,
scheme: UpstreamProxyScheme,
@@ -50,6 +50,20 @@ pub(crate) struct UpstreamProxyConfig {
password: Option<String>,
}
impl std::fmt::Debug for UpstreamProxyConfig {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("UpstreamProxyConfig")
.field("url", &self.redacted_url())
.field("scheme", &self.scheme)
.field("host", &self.host)
.field("port", &self.port)
.field("username", &self.username.as_ref().map(|_| "[REDACTED]"))
.field("password", &self.password.as_ref().map(|_| "[REDACTED]"))
.finish()
}
}
impl UpstreamProxyConfig {
pub(crate) fn parse(raw: &str) -> Result<Self, String> {
let trimmed = raw.trim();
@@ -75,6 +89,18 @@ impl UpstreamProxyConfig {
.filter(|value| !value.is_empty())
.ok_or_else(|| "upstream proxy URL must include a host".to_string())?
.to_string();
// A proxy URL identifies the proxy origin. Path/query/fragment
// components are not part of HTTP CONNECT or SOCKS negotiation and
// silently ignoring them can make the configured endpoint differ
// from what operators see in configuration and logs.
if !matches!(parsed.path(), "" | "/")
|| parsed.query().is_some()
|| parsed.fragment().is_some()
{
return Err(
"upstream proxy URL must be an origin without path, query, or fragment".to_string(),
);
}
let port = parsed.port().unwrap_or(match scheme {
UpstreamProxyScheme::Http => 80,
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => 1080,
@@ -181,6 +207,38 @@ pub(crate) async fn connect_target_via_proxy(
Ok(tcp)
}
pub(crate) async fn connect_validated_target_via_proxy(
proxy: &UpstreamProxyConfig,
target_addr: SocketAddr,
options: ProxyConnectOptions,
) -> io::Result<TcpStream> {
let mut tcp = connect_proxy_tcp(
proxy,
options.connect_timeout,
options.tcp_nodelay,
options.tcp_keepalive,
options.ip_family,
)
.await?;
match proxy.scheme() {
UpstreamProxyScheme::Http => {
http_connect(&mut tcp, &target_addr.to_string(), proxy).await?;
}
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => {
socks5_connect(
&mut tcp,
proxy,
&target_addr.ip().to_string(),
target_addr.port(),
)
.await?;
}
}
Ok(tcp)
}
pub(crate) async fn connect_proxy_tcp(
proxy: &UpstreamProxyConfig,
connect_timeout: Duration,
@@ -188,16 +246,19 @@ pub(crate) async fn connect_proxy_tcp(
tcp_keepalive: Option<Duration>,
ip_family: IpFamily,
) -> io::Result<TcpStream> {
let resolved = tokio::time::timeout(
connect_timeout,
tokio::net::lookup_host((proxy.host(), proxy.port())),
)
.await
.map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "proxy DNS timeout"))?
.map_err(|err| io::Error::other(format!("proxy DNS failed: {err}")))?;
let resolved =
aether_http::lookup_host_with_limits(proxy.host(), proxy.port(), connect_timeout)
.await
.map_err(|err| {
if err.kind() == io::ErrorKind::TimedOut {
io::Error::new(io::ErrorKind::TimedOut, "proxy DNS timeout")
} else {
io::Error::other(format!("proxy DNS failed: {err}"))
}
})?;
let mut last_error = None;
for addr in resolved.filter(|addr| ip_family.allows(*addr)) {
for addr in resolved.into_iter().filter(|addr| ip_family.allows(*addr)) {
match tokio::time::timeout(connect_timeout, TcpStream::connect(addr)).await {
Ok(Ok(stream)) => {
configure_tcp_stream(&stream, tcp_nodelay, tcp_keepalive)?;
@@ -395,12 +456,16 @@ pub(crate) async fn socks5_target_address(
request.push(host.len() as u8);
request.extend_from_slice(host);
} else {
let mut resolved = tokio::net::lookup_host((target_host, target_port))
.await
.map_err(|err| io::Error::other(format!("SOCKS5 target DNS failed: {err}")))?;
let addr = resolved
.next()
.ok_or_else(|| io::Error::other("SOCKS5 target DNS returned no addresses"))?;
let addr = aether_http::lookup_host_with_limits(
target_host,
target_port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|err| io::Error::other(format!("SOCKS5 target DNS failed: {err}")))?
.into_iter()
.next()
.ok_or_else(|| io::Error::other("SOCKS5 target DNS returned no addresses"))?;
push_socks5_socket_address(&mut request, addr);
}
request.extend_from_slice(&target_port.to_be_bytes());
@@ -481,6 +546,13 @@ mod tests {
proxy.basic_auth_header().as_deref(),
Some("Basic dXNlcjpwYXNz")
);
let debug = format!("{proxy:?}");
assert!(!debug.contains("user:pass"));
assert!(!debug.contains("Some(\"user\")"));
assert!(!debug.contains("Some(\"pass\")"));
assert!(!debug.contains("dXNlcjpwYXNz"));
assert!(debug.contains("[REDACTED]"));
}
#[test]
@@ -490,4 +562,24 @@ mod tests {
assert!(error.contains("unsupported upstream proxy scheme"));
}
#[test]
fn rejects_proxy_urls_with_non_origin_components() {
for value in [
"http://proxy.example/path",
"http://proxy.example?token=secret",
"socks5://proxy.example#fragment",
] {
let error = UpstreamProxyConfig::parse(value)
.expect_err("proxy URL with non-origin components should be rejected");
assert!(
error.contains("without path, query, or fragment"),
"unexpected error for {value}: {error}"
);
}
for value in ["http://proxy.example", "http://proxy.example/"] {
UpstreamProxyConfig::parse(value).expect("root proxy origin should be accepted");
}
}
}
+2 -2
View File
@@ -226,7 +226,7 @@ mod tests {
let (config, tunnel_security) = parse_config_and_security(&[
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
@@ -252,7 +252,7 @@ mod tests {
let (config, tunnel_security) = parse_config_and_security(&[
"aether-tunnel",
"--aether-url",
"http://example.com",
"http://127.0.0.1:8084",
"--management-token",
"ae_test",
"--node-name",
+131 -17
View File
@@ -2,8 +2,15 @@
//!
//! These are standalone helpers not tied to any specific client or service.
use aether_http::{build_http_client, HttpClientConfig};
use std::net::IpAddr;
use aether_http::{
build_http_client, is_private_or_reserved_ip, read_response_bytes_with_limit, HttpClientConfig,
};
use tracing::{debug, info};
use url::Url;
const MAX_NETWORK_DISCOVERY_RESPONSE_BYTES: usize = 16 * 1024;
/// Auto-detect public IP by querying external services.
pub async fn detect_public_ip() -> anyhow::Result<String> {
@@ -22,10 +29,21 @@ pub async fn detect_public_ip() -> anyhow::Result<String> {
for endpoint in &endpoints {
match client.get(*endpoint).send().await {
Ok(resp) if resp.status().is_success() => {
let ip = resp.text().await?.trim().to_string();
if !ip.is_empty() {
info!(ip = %ip, source = %endpoint, "detected public IP");
return Ok(ip);
match read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES)
.await
{
Ok(body) => {
let text = String::from_utf8_lossy(&body);
if let Some(ip) = parse_ip(text.as_ref()) {
let ip = ip.to_string();
info!(ip = %ip, source = %endpoint, "detected public IP");
return Ok(ip);
}
debug!(endpoint = %endpoint, "IP detection response was not a valid IP");
}
Err(error) => {
debug!(endpoint = %endpoint, error = %error, "IP detection response rejected");
}
}
}
Ok(resp) => {
@@ -47,8 +65,20 @@ pub async fn detect_public_ip() -> anyhow::Result<String> {
/// This is best-effort and non-sensitive -- region detection should never
/// block startup.
pub async fn detect_region(ip: &str) -> Option<String> {
// The value may come directly from the command line/environment. Parse
// it before putting it into a URL and avoid disclosing private or reserved
// addresses to third-party geolocation services. We intentionally do not
// reject such values from registration: controlled internal deployments
// can still advertise their configured address and region explicitly.
let ip = parse_ip(ip)?;
if is_private_or_reserved_ip(ip) {
debug!(ip = %ip, "skipping region detection for private or reserved IP");
return None;
}
let ip = ip.to_string();
// Try HTTPS provider first
let https_url = format!("https://ipinfo.io/{}/country", ip);
let https_url = ipinfo_url(&ip)?;
let client = build_http_client(&HttpClientConfig {
request_timeout_ms: Some(5_000),
@@ -60,27 +90,29 @@ pub async fn detect_region(ip: &str) -> Option<String> {
// Try ipinfo.io (HTTPS, returns plain text country code)
if let Ok(resp) = client.get(&https_url).send().await {
if resp.status().is_success() {
if let Ok(text) = resp.text().await {
let code = text.trim();
if !code.is_empty() && code.len() <= 3 {
if let Ok(body) =
read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES).await
{
let text = String::from_utf8_lossy(&body);
if let Some(code) = normalize_country_code(text.as_ref()) {
info!(region = %code, ip = %ip, source = "ipinfo.io", "detected region");
return Some(code.to_string());
return Some(code);
}
}
}
}
// Fallback: ip-api.com (HTTP only on free tier, non-sensitive data)
let http_url = format!("http://ip-api.com/json/{}?fields=countryCode", ip);
let http_url = ip_api_url(&ip)?;
match client.get(&http_url).send().await {
Ok(resp) if resp.status().is_success() => {
let body: serde_json::Value = resp.json().await.ok()?;
let code = body.get("countryCode")?.as_str()?;
if code.is_empty() {
return None;
}
let body = read_response_bytes_with_limit(resp, MAX_NETWORK_DISCOVERY_RESPONSE_BYTES)
.await
.ok()?;
let body: serde_json::Value = serde_json::from_slice(&body).ok()?;
let code = normalize_country_code(body.get("countryCode")?.as_str()?)?;
info!(region = %code, ip = %ip, source = "ip-api.com", "detected region");
Some(code.to_string())
Some(code)
}
_ => {
debug!(ip = %ip, "region detection failed");
@@ -88,3 +120,85 @@ pub async fn detect_region(ip: &str) -> Option<String> {
}
}
}
fn parse_ip(value: &str) -> Option<IpAddr> {
let value = value.trim();
if value.is_empty() || value.len() > 45 {
return None;
}
value.parse().ok()
}
fn normalize_country_code(value: &str) -> Option<String> {
let value = value.trim();
if !(2..=3).contains(&value.len()) || !value.bytes().all(|byte| byte.is_ascii_alphabetic()) {
return None;
}
Some(value.to_ascii_uppercase())
}
fn ipinfo_url(ip: &str) -> Option<String> {
let mut url = Url::parse("https://ipinfo.io").ok()?;
url.path_segments_mut().ok()?.push(ip).push("country");
Some(url.into())
}
fn ip_api_url(ip: &str) -> Option<String> {
let mut url = Url::parse("http://ip-api.com").ok()?;
url.path_segments_mut().ok()?.push("json").push(ip);
url.query_pairs_mut().append_pair("fields", "countryCode");
Some(url.into())
}
#[cfg(test)]
mod tests {
use super::{ip_api_url, ipinfo_url, normalize_country_code, parse_ip};
#[test]
fn parses_only_bounded_ip_values() {
assert_eq!(parse_ip(" 8.8.8.8\n"), Some("8.8.8.8".parse().unwrap()));
assert_eq!(
parse_ip("2001:4860:4860::8888"),
Some("2001:4860:4860::8888".parse().unwrap())
);
for value in ["", "8.8.8.8?x=1", "8.8.8.8\nX-Injected: yes", "not-an-ip"] {
assert!(parse_ip(value).is_none(), "accepted invalid IP: {value:?}");
}
}
#[test]
fn builds_urls_without_allowing_path_or_query_injection() {
let ip = "2001:4860:4860::8888";
let info = ipinfo_url(ip).expect("fixed URL should parse");
let info = url::Url::parse(&info).unwrap();
assert_eq!(info.host_str(), Some("ipinfo.io"));
assert_eq!(
info.path_segments().unwrap().collect::<Vec<_>>(),
["2001:4860:4860::8888", "country"]
);
assert_eq!(info.query(), None);
assert_eq!(info.fragment(), None);
let api = ip_api_url(ip).expect("fixed URL should parse");
let api = url::Url::parse(&api).unwrap();
assert_eq!(api.host_str(), Some("ip-api.com"));
assert_eq!(api.query(), Some("fields=countryCode"));
assert_eq!(
api.path_segments().unwrap().collect::<Vec<_>>(),
["json", "2001:4860:4860::8888"]
);
assert_eq!(api.fragment(), None);
}
#[test]
fn country_codes_are_ascii_and_normalized() {
assert_eq!(normalize_country_code(" us\n"), Some("US".to_string()));
assert_eq!(normalize_country_code("GBR"), Some("GBR".to_string()));
for value in ["", "U", "US!", "US\nX", "中国"] {
assert!(
normalize_country_code(value).is_none(),
"accepted {value:?}"
);
}
}
}
+117 -18
View File
@@ -1,13 +1,21 @@
use aether_http::{build_http_client, jittered_delay_for_retry, HttpClientConfig, HttpRetryConfig};
use aether_http::{
apply_http_client_config, jittered_delay_for_retry, read_response_bytes_with_limit,
HttpClientConfig, HttpRetryConfig, ResponseBodyReadError,
};
use aether_runtime::summarize_text_payload;
use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use tokio::time::sleep;
use tracing::{debug, error, info};
use crate::config::{effective_tunnel_security, Config, ServerEntry, TunnelSecurity};
use crate::config::{
aether_url_for_log, effective_tunnel_security, validate_aether_url, Config, ServerEntry,
TunnelSecurity,
};
use crate::hardware::HardwareInfo;
const MAX_AETHER_CONTROL_RESPONSE_BYTES: usize = 256 * 1024;
#[derive(Debug, Serialize)]
struct RegisterRequest {
name: String,
@@ -32,6 +40,7 @@ struct RegisterRequest {
#[derive(Debug, Deserialize)]
pub struct RegisterResponse {
pub node_id: String,
pub tunnel_generation: String,
}
/// Remote configuration pushed by the Aether management backend.
@@ -58,7 +67,7 @@ pub struct AetherClient {
impl AetherClient {
pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self {
let http = build_http_client(&HttpClientConfig {
let client_config = HttpClientConfig {
connect_timeout_ms: Some(config.aether_connect_timeout_secs.saturating_mul(1_000)),
request_timeout_ms: Some(config.aether_request_timeout_secs.saturating_mul(1_000)),
pool_idle_timeout_ms: Some(config.aether_pool_idle_timeout_secs.saturating_mul(1_000)),
@@ -75,8 +84,18 @@ impl AetherClient {
.effective_aether_outbound_proxy_url()
.map(str::to_string),
..HttpClientConfig::default()
})
.expect("failed to create HTTP client");
};
let mut builder = apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&client_config,
);
if let Some(proxy_url) = client_config.proxy_url.as_deref() {
builder = builder
.proxy(reqwest::Proxy::all(proxy_url).expect("invalid Aether outbound proxy URL"));
}
let http = builder.build().expect("failed to create HTTP client");
let retry = HttpRetryConfig {
max_attempts: config.aether_retry_max_attempts,
@@ -87,7 +106,7 @@ impl AetherClient {
Self {
http,
base_url: aether_url.trim_end_matches('/').to_string(),
base_url: aether_url.trim().trim_end_matches('/').to_string(),
token: management_token.to_string(),
retry,
}
@@ -95,7 +114,7 @@ impl AetherClient {
/// Register this node with Aether (idempotent upsert by ip:port).
///
/// Returns the stable node_id assigned by Aether.
/// Returns the stable node identity and the current server-issued tunnel generation.
pub async fn register(
&self,
config: &Config,
@@ -103,7 +122,8 @@ impl AetherClient {
node_name: &str,
public_ip: &str,
hw: Option<&HardwareInfo>,
) -> anyhow::Result<String> {
) -> anyhow::Result<RegisterResponse> {
validate_aether_url(&self.base_url)?;
let url = format!("{}/api/admin/proxy-nodes/register", self.base_url);
let effective_security = effective_tunnel_security(
&server.aether_url,
@@ -130,7 +150,7 @@ impl AetherClient {
};
info!(
url = %url,
url = %aether_url_for_log(&url),
name = %body.name,
ip = %body.ip,
"registering with Aether"
@@ -146,11 +166,22 @@ impl AetherClient {
},
"register",
)
.await?;
.await
.map_err(|error| {
anyhow::anyhow!("register request failed ({})", reqwest_error_kind(&error))
})?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
let body = read_response_bytes_with_limit(resp, MAX_AETHER_CONTROL_RESPONSE_BYTES)
.await
.map_err(|error| {
anyhow::anyhow!(
"register failed (HTTP {status}): {}",
response_body_error_kind(&error)
)
})?;
let text = String::from_utf8_lossy(&body);
let summary = summarize_text_payload(&text);
anyhow::bail!(
"register failed (HTTP {}): response body redacted (bytes={}, sha256={})",
@@ -160,13 +191,23 @@ impl AetherClient {
);
}
let data: RegisterResponse = resp.json().await?;
info!(node_id = %data.node_id, "registered successfully");
Ok(data.node_id)
let body = read_response_bytes_with_limit(resp, MAX_AETHER_CONTROL_RESPONSE_BYTES)
.await
.map_err(|error| {
anyhow::anyhow!(
"failed to read Aether registration response ({})",
response_body_error_kind(&error)
)
})?;
let data: RegisterResponse = serde_json::from_slice(&body)?;
validate_registration_binding(&data)?;
info!(node_id = %data.node_id, tunnel_generation = %data.tunnel_generation, "registered successfully");
Ok(data)
}
/// Unregister this node from Aether (graceful shutdown).
pub async fn unregister(&self, node_id: &str) -> anyhow::Result<()> {
validate_aether_url(&self.base_url)?;
let url = format!("{}/api/admin/proxy-nodes/unregister", self.base_url);
let body = UnregisterRequest {
node_id: node_id.to_string(),
@@ -193,7 +234,15 @@ impl AetherClient {
}
Ok(r) => {
let status = r.status();
let text = r.text().await.unwrap_or_default();
let body = read_response_bytes_with_limit(r, MAX_AETHER_CONTROL_RESPONSE_BYTES)
.await
.map_err(|error| {
anyhow::anyhow!(
"unregister failed (HTTP {status}): {}",
response_body_error_kind(&error)
)
})?;
let text = String::from_utf8_lossy(&body);
let summary = summarize_text_payload(&text);
error!(
status = %status,
@@ -210,8 +259,9 @@ impl AetherClient {
}
Err(e) => {
// Best-effort during shutdown
error!(error = %e, "unregister request failed");
anyhow::bail!("unregister request failed: {}", e);
let kind = reqwest_error_kind(&e);
error!(error_kind = kind, "unregister request failed");
anyhow::bail!("unregister request failed ({kind})");
}
}
}
@@ -250,7 +300,7 @@ impl AetherClient {
let sleep_for = jittered_delay_for_retry(self.retry, attempt - 1);
debug!(
attempt,
error = %e,
error_kind = reqwest_error_kind(&e),
sleep_ms = sleep_for.as_millis(),
label,
"Aether request retrying"
@@ -265,6 +315,55 @@ impl AetherClient {
}
}
fn validate_registration_binding(response: &RegisterResponse) -> anyhow::Result<()> {
for (field, value) in [
("node_id", response.node_id.as_str()),
("tunnel_generation", response.tunnel_generation.as_str()),
] {
if value.is_empty()
|| value.len() > 128
|| !value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
{
anyhow::bail!("registration response contains an invalid {field}");
}
}
Ok(())
}
/// Reqwest error messages can include the complete request URL. Registration
/// URLs are derived from configuration, whose path may contain an operator
/// secret, so expose only a stable transport category at this boundary.
fn 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_redirect() {
"redirect"
} else if error.is_body() {
"body"
} else if error.is_decode() {
"decode"
} else {
"transport"
}
}
fn response_body_error_kind(error: &ResponseBodyReadError) -> String {
match error {
ResponseBodyReadError::TooLarge { max_bytes } => {
format!("response too large (max {max_bytes} bytes)")
}
ResponseBodyReadError::Read(error) => {
format!("response read failed ({})", reqwest_error_kind(error))
}
}
}
fn should_retry_status(status: StatusCode) -> bool {
status.is_server_error()
|| status == StatusCode::TOO_MANY_REQUESTS
-66
View File
@@ -1,66 +0,0 @@
//! Safe DNS resolver for reqwest that reuses validated addresses from DnsCache.
//!
//! This resolver ensures reqwest connects only to addresses that have been
//! previously validated by `target_filter::validate_target()`, eliminating
//! the TOCTTOU gap where DNS rebinding could redirect traffic to private IPs.
use std::net::SocketAddr;
use std::sync::Arc;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use crate::target_filter::{self, DnsCache};
/// A DNS resolver that serves validated public addresses from the shared DnsCache.
///
/// When reqwest needs to resolve a hostname, this resolver returns addresses
/// from the cache (populated by `validate_target()` during request validation).
/// If the hostname is not in cache (shouldn't happen in normal flow), it
/// performs a fresh resolution with private-IP filtering.
pub struct SafeDnsResolver {
dns_cache: Arc<DnsCache>,
}
impl SafeDnsResolver {
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
Self { dns_cache }
}
}
impl Resolve for SafeDnsResolver {
fn resolve(&self, name: Name) -> Resolving {
let dns_cache = Arc::clone(&self.dns_cache);
Box::pin(async move {
let host = name.as_str();
// Try cache first (should be populated by validate_target).
// reqwest resolves by hostname only (no port), so use host-only lookup.
if let Some(addrs) = dns_cache.get_by_host(host).await {
let socket_addrs: Vec<SocketAddr> = (*addrs).clone();
return Ok(Box::new(socket_addrs.into_iter()) as Addrs);
}
// Fallback: resolve with private-IP filtering (defensive).
// This path should rarely be hit since validate_target() runs first.
// We don't know the real port here (reqwest Resolve only gives hostname),
// so resolve directly without caching to avoid polluting the cache with
// an incorrect port-based key.
let addr_str = format!("{}:0", host);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { Box::new(e) })?
.filter(|addr| !target_filter::is_private_ip(&addr.ip()))
.collect();
if resolved.is_empty() {
return Err(Box::new(std::io::Error::other(format!(
"all resolved addresses for {} are private/reserved",
host
)))
as Box<dyn std::error::Error + Send + Sync>);
}
Ok(Box::new(resolved.into_iter()) as Addrs)
})
}
}
+511 -56
View File
@@ -3,8 +3,11 @@
//! Supports the host-native service manager we currently target:
//! `systemd` on most Linux distributions and `OpenRC` on Alpine.
#[cfg(unix)]
use std::fs::OpenOptions;
use std::io::ErrorKind;
#[cfg(unix)]
use std::io::Write;
use std::path::Path;
use std::process::{Command, ExitStatus, Stdio};
@@ -18,6 +21,16 @@ const OPENRC_LOG_DIR: &str = "/var/log/aether-tunnel";
const OPENRC_STDOUT_LOG: &str = "/var/log/aether-tunnel/current.log";
const OPENRC_STDERR_LOG: &str = "/var/log/aether-tunnel/error.log";
const SYSTEMCTL_BINS: &[&str] = &[
"/usr/bin/systemctl",
"/bin/systemctl",
"/run/current-system/sw/bin/systemctl",
];
const JOURNALCTL_BINS: &[&str] = &[
"/usr/bin/journalctl",
"/bin/journalctl",
"/run/current-system/sw/bin/journalctl",
];
const OPENRC_RUN_BINS: &[&str] = &["/sbin/openrc-run", "/usr/sbin/openrc-run", "openrc-run"];
const OPENRC_SERVICE_BINS: &[&str] = &["/sbin/rc-service", "/usr/sbin/rc-service", "rc-service"];
const OPENRC_UPDATE_BINS: &[&str] = &["/sbin/rc-update", "/usr/sbin/rc-update", "rc-update"];
@@ -143,7 +156,7 @@ pub fn cmd_logs() -> anyhow::Result<()> {
ensure_openrc_logs_readable()?;
}
let status = match manager {
ServiceManager::Systemd => Command::new("journalctl")
ServiceManager::Systemd => Command::new(journalctl_bin())
.args(["-u", SERVICE_NAME, "-f", "--no-pager", "-n", "100"])
.status()?,
ServiceManager::OpenRc => Command::new(tail_bin())
@@ -281,9 +294,15 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> {
.to_str()
.unwrap_or("/");
validate_service_unit_path(exe_str, "binary")?;
validate_service_unit_path(config_str, "config")?;
validate_service_unit_path(working_dir, "working directory")?;
validate_root_managed_service_file(&exe_path, "binary", false)?;
validate_root_managed_service_file(&config_abs, "config", true)?;
if Path::new(SYSTEMD_UNIT_PATH).exists() {
eprintln!(" Stopping existing service...");
let _ = Command::new("systemctl")
let _ = Command::new(systemctl_bin())
.args(["stop", SERVICE_NAME])
.status();
}
@@ -293,34 +312,12 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> {
eprintln!(" Config: {}", config_str);
eprintln!(" WorkDir: {}", working_dir);
let unit_content = format!(
"[Unit]\n\
Description=Aether Tunnel\n\
After=network.target\n\
\n\
[Service]\n\
Type=simple\n\
WorkingDirectory={working_dir}\n\
Environment=AETHER_TUNNEL_CONFIG={config_str}\n\
Environment=AETHER_TUNNEL_SERVICE_MANAGER=systemd\n\
Environment=AETHER_TUNNEL_LOG_DESTINATION=both\n\
Environment=AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel\n\
ExecStart={exe_str}\n\
Restart=on-failure\n\
RestartSec=5\n\
LimitNOFILE=65535\n\
UMask=0077\n\
LogsDirectory=aether-tunnel\n\
LogsDirectoryMode=0750\n\
\n\
[Install]\n\
WantedBy=multi-user.target\n",
);
std::fs::write(SYSTEMD_UNIT_PATH, &unit_content)?;
let unit_content = render_systemd_unit(exe_str, config_str, working_dir)?;
write_service_definition(SYSTEMD_UNIT_PATH, &unit_content, 0o644)?;
eprintln!(" Enabling and starting service...");
run_cmd("systemctl", &["daemon-reload"])?;
run_cmd("systemctl", &["enable", "--now", SERVICE_NAME])?;
run_cmd(systemctl_bin(), &["daemon-reload"])?;
run_cmd(systemctl_bin(), &["enable", "--now", SERVICE_NAME])?;
eprintln!();
if manager_is_active(ServiceManager::Systemd) {
@@ -333,6 +330,43 @@ fn install_systemd_service(config_path: &Path) -> anyhow::Result<()> {
Ok(())
}
fn render_systemd_unit(
exe_path: &str,
config_path: &str,
working_dir: &str,
) -> anyhow::Result<String> {
validate_service_unit_path(exe_path, "binary")?;
validate_service_unit_path(config_path, "config")?;
validate_service_unit_path(working_dir, "working directory")?;
let exe_path = systemd_quote(exe_path);
let working_dir = systemd_quote(working_dir);
let config_env = systemd_quote(&format!("AETHER_TUNNEL_CONFIG={config_path}"));
Ok(format!(
"[Unit]\n\
Description=Aether Tunnel\n\
After=network.target\n\
\n\
[Service]\n\
Type=simple\n\
WorkingDirectory={working_dir}\n\
Environment={config_env}\n\
Environment=AETHER_TUNNEL_SERVICE_MANAGER=systemd\n\
Environment=AETHER_TUNNEL_LOG_DESTINATION=both\n\
Environment=AETHER_TUNNEL_LOG_DIR=/var/log/aether-tunnel\n\
ExecStart={exe_path}\n\
Restart=on-failure\n\
RestartSec=5\n\
LimitNOFILE=65535\n\
UMask=0077\n\
LogsDirectory=aether-tunnel\n\
LogsDirectoryMode=0750\n\
\n\
[Install]\n\
WantedBy=multi-user.target\n",
))
}
fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> {
let exe_path = std::env::current_exe()?.canonicalize()?;
let exe_str = exe_path
@@ -350,6 +384,12 @@ fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> {
.to_str()
.unwrap_or("/");
validate_service_unit_path(exe_str, "binary")?;
validate_service_unit_path(config_str, "config")?;
validate_service_unit_path(working_dir, "working directory")?;
validate_root_managed_service_file(&exe_path, "binary", false)?;
validate_root_managed_service_file(&config_abs, "config", true)?;
if Path::new(OPENRC_INIT_PATH).exists() {
eprintln!(" Stopping existing service...");
let _ = Command::new(openrc_service_bin())
@@ -357,12 +397,9 @@ fn install_openrc_service(config_path: &Path) -> anyhow::Result<()> {
.status();
}
std::fs::create_dir_all(OPENRC_LOG_DIR)?;
touch_log(OPENRC_STDOUT_LOG)?;
touch_log(OPENRC_STDERR_LOG)?;
set_mode(OPENRC_LOG_DIR, 0o750)?;
set_mode(OPENRC_STDOUT_LOG, 0o640)?;
set_mode(OPENRC_STDERR_LOG, 0o640)?;
ensure_private_service_directory(Path::new(OPENRC_LOG_DIR), 0o750)?;
open_private_service_log(Path::new(OPENRC_STDOUT_LOG), 0o640)?;
open_private_service_log(Path::new(OPENRC_STDERR_LOG), 0o640)?;
eprintln!(" Generating OpenRC init script...");
eprintln!(" Binary: {}", exe_str);
@@ -439,8 +476,7 @@ stop() {{
shell_quote("AETHER_TUNNEL_LOG_DESTINATION=both"),
shell_quote(&format!("AETHER_TUNNEL_LOG_DIR={OPENRC_LOG_DIR}")),
);
std::fs::write(OPENRC_INIT_PATH, &init_content)?;
set_mode(OPENRC_INIT_PATH, 0o755)?;
write_service_definition(OPENRC_INIT_PATH, &init_content, 0o755)?;
eprintln!(" Enabling and starting service...");
run_cmd(openrc_update_bin(), &["add", SERVICE_NAME, "default"])?;
@@ -459,7 +495,7 @@ stop() {{
fn uninstall_systemd_service() -> anyhow::Result<()> {
eprintln!(" Stopping and removing existing service...");
let _ = Command::new("systemctl")
let _ = Command::new(systemctl_bin())
.args(["disable", "--now", SERVICE_NAME])
.status();
@@ -468,7 +504,7 @@ fn uninstall_systemd_service() -> anyhow::Result<()> {
eprintln!(" Removed {}", SYSTEMD_UNIT_PATH);
}
run_cmd("systemctl", &["daemon-reload"])?;
run_cmd(systemctl_bin(), &["daemon-reload"])?;
eprintln!(" Service uninstalled.");
Ok(())
}
@@ -493,28 +529,28 @@ fn uninstall_openrc_service() -> anyhow::Result<()> {
fn start_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["start", SERVICE_NAME]),
ServiceManager::Systemd => run_cmd(systemctl_bin(), &["start", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "start"]),
}
}
fn stop_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["stop", SERVICE_NAME]),
ServiceManager::Systemd => run_cmd(systemctl_bin(), &["stop", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "stop"]),
}
}
fn restart_manager(manager: ServiceManager) -> anyhow::Result<()> {
match manager {
ServiceManager::Systemd => run_cmd("systemctl", &["restart", SERVICE_NAME]),
ServiceManager::Systemd => run_cmd(systemctl_bin(), &["restart", SERVICE_NAME]),
ServiceManager::OpenRc => run_cmd(openrc_service_bin(), &[SERVICE_NAME, "restart"]),
}
}
fn manager_status(manager: ServiceManager) -> anyhow::Result<ExitStatus> {
let status = match manager {
ServiceManager::Systemd => Command::new("systemctl")
ServiceManager::Systemd => Command::new(systemctl_bin())
.args(["status", SERVICE_NAME])
.status()?,
ServiceManager::OpenRc => Command::new(openrc_service_bin())
@@ -528,7 +564,7 @@ fn manager_is_active(manager: ServiceManager) -> bool {
match manager {
ServiceManager::Systemd => {
Path::new(SYSTEMD_UNIT_PATH).exists()
&& Command::new("systemctl")
&& Command::new(systemctl_bin())
.args(["is-active", "--quiet", SERVICE_NAME])
.stdout(Stdio::null())
.stderr(Stdio::null())
@@ -562,7 +598,8 @@ fn print_post_install_commands() {
fn is_systemd_available() -> bool {
Path::new("/run/systemd/system").exists()
&& Command::new("systemctl")
&& has_absolute_candidate(SYSTEMCTL_BINS)
&& Command::new(systemctl_bin())
.arg("--version")
.stdout(Stdio::null())
.stderr(Stdio::null())
@@ -610,28 +647,446 @@ fn pick_bin(candidates: &[&'static str]) -> &'static str {
.iter()
.copied()
.find(|candidate| candidate.starts_with('/') && Path::new(candidate).exists())
.unwrap_or_else(|| candidates[candidates.len() - 1])
.or_else(|| {
candidates
.iter()
.copied()
.find(|candidate| candidate.starts_with('/'))
})
.expect("trusted binary candidate list must include an absolute path")
}
fn systemctl_bin() -> &'static str {
pick_bin(SYSTEMCTL_BINS)
}
fn journalctl_bin() -> &'static str {
pick_bin(JOURNALCTL_BINS)
}
fn validate_service_unit_path(value: &str, label: &str) -> anyhow::Result<()> {
if !Path::new(value).is_absolute() {
anyhow::bail!("{} path must be absolute", label);
}
if value.chars().any(char::is_control) || value.contains(['%', '$']) {
anyhow::bail!(
"{} path contains control or service-manager expansion characters",
label
);
}
Ok(())
}
fn validate_root_managed_service_file(
path: &Path,
label: &str,
require_private_file: bool,
) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::MetadataExt;
let file_metadata = std::fs::symlink_metadata(path)?;
if !file_metadata.is_file() || file_metadata.file_type().is_symlink() {
anyhow::bail!("service {} must be a regular non-symlink file", label);
}
if file_metadata.uid() != 0 || file_metadata.mode() & 0o022 != 0 {
anyhow::bail!(
"service {} must be owned by root and not writable by group or other users; use the official installer or move it to a root-managed path",
label
);
}
if require_private_file && (file_metadata.mode() & 0o077 != 0 || file_metadata.nlink() != 1)
{
anyhow::bail!(
"service {} contains credentials and must be owner-only with exactly one hard link",
label
);
}
let mut ancestor = path.parent();
while let Some(directory) = ancestor {
let metadata = std::fs::symlink_metadata(directory)?;
if !metadata.is_dir()
|| metadata.file_type().is_symlink()
|| metadata.uid() != 0
|| metadata.mode() & 0o022 != 0
{
anyhow::bail!(
"service {} parent '{}' must be a root-owned directory that is not writable by group or other users",
label,
directory.display()
);
}
ancestor = directory.parent();
}
Ok(())
}
#[cfg(not(unix))]
{
let _ = (path, label, require_private_file);
anyhow::bail!("managed tunnel services require Unix ownership checks")
}
}
fn systemd_quote(value: &str) -> String {
format!("\"{}\"", value.replace('\\', "\\\\").replace('\"', "\\\""))
}
fn write_service_definition(path: &str, content: &str, mode: u32) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt};
let requested_path = Path::new(path);
let requested_parent = requested_path
.parent()
.ok_or_else(|| anyhow::anyhow!("service definition path has no parent"))?;
let file_name = requested_path
.file_name()
.ok_or_else(|| anyhow::anyhow!("service definition path has no file name"))?;
let parent = std::fs::canonicalize(requested_parent)?;
let path = parent.join(file_name);
validate_private_service_directory(&parent)?;
validate_replaceable_service_file(&path)?;
let temporary = parent.join(format!(
".aether-tunnel-service-{}-{}.tmp",
std::process::id(),
uuid::Uuid::new_v4()
));
let mut options = OpenOptions::new();
options
.write(true)
.create_new(true)
.mode(mode)
.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW);
let mut file = options.open(&temporary)?;
let result = (|| -> anyhow::Result<()> {
// SAFETY: geteuid has no preconditions and does not retain pointers.
let effective_uid = unsafe { libc::geteuid() };
let metadata = file.metadata()?;
if !metadata.is_file() || metadata.uid() != effective_uid || metadata.nlink() != 1 {
anyhow::bail!("temporary service definition has unsafe ownership or links");
}
file.set_permissions(std::fs::Permissions::from_mode(mode))?;
file.write_all(content.as_bytes())?;
file.sync_all()?;
drop(file);
std::fs::rename(&temporary, &path)?;
std::fs::File::open(&parent)?.sync_all()?;
Ok(())
})();
if result.is_err() {
let _ = std::fs::remove_file(&temporary);
}
return result;
}
#[cfg(not(unix))]
{
let _ = (path, content, mode);
anyhow::bail!("managed service definitions require Unix filesystem checks")
}
}
fn shell_quote(value: &str) -> String {
format!("'{}'", value.replace('\'', "'\"'\"'"))
}
fn touch_log(path: &str) -> anyhow::Result<()> {
OpenOptions::new().create(true).append(true).open(path)?;
Ok(())
}
fn set_mode(path: &str, mode: u32) -> anyhow::Result<()> {
fn validate_private_service_directory(path: &Path) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let perms = std::fs::Permissions::from_mode(mode);
std::fs::set_permissions(path, perms)?;
use std::os::unix::fs::MetadataExt;
// SAFETY: geteuid has no preconditions and does not retain pointers.
let effective_uid = unsafe { libc::geteuid() };
let metadata = std::fs::symlink_metadata(path)?;
if metadata.file_type().is_symlink()
|| !metadata.is_dir()
|| (metadata.uid() != effective_uid && metadata.uid() != 0)
|| metadata.mode() & 0o022 != 0
{
anyhow::bail!(
"service directory '{}' has unsafe ownership or permissions",
path.display()
);
}
let canonical = std::fs::canonicalize(path)?;
let mut ancestor = canonical.parent();
while let Some(directory) = ancestor {
let metadata = std::fs::symlink_metadata(directory)?;
let mode = metadata.mode();
if metadata.file_type().is_symlink()
|| !metadata.is_dir()
|| (metadata.uid() != effective_uid && metadata.uid() != 0)
|| (mode & 0o022 != 0 && mode & 0o1000 == 0)
{
anyhow::bail!(
"service directory ancestor '{}' has unsafe ownership or permissions",
directory.display()
);
}
ancestor = directory.parent();
}
return Ok(());
}
#[cfg(not(unix))]
let _ = (path, mode);
{
let _ = path;
anyhow::bail!("managed service directories require Unix filesystem checks")
}
}
Ok(())
fn validate_replaceable_service_file(path: &Path) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::MetadataExt;
match std::fs::symlink_metadata(path) {
Ok(metadata) => {
// SAFETY: geteuid has no preconditions and does not retain pointers.
let effective_uid = unsafe { libc::geteuid() };
if metadata.file_type().is_symlink()
|| !metadata.is_file()
|| metadata.uid() != effective_uid
|| metadata.nlink() != 1
{
anyhow::bail!(
"service file '{}' must be a regular single-link file owned by the current user",
path.display()
);
}
}
Err(error) if error.kind() == ErrorKind::NotFound => {}
Err(error) => return Err(error.into()),
}
return Ok(());
}
#[cfg(not(unix))]
{
let _ = path;
anyhow::bail!("managed service files require Unix filesystem checks")
}
}
fn ensure_private_service_directory(path: &Path, mode: u32) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::{DirBuilderExt, MetadataExt, OpenOptionsExt, PermissionsExt};
let requested_parent = path
.parent()
.ok_or_else(|| anyhow::anyhow!("service directory has no parent"))?;
let file_name = path
.file_name()
.ok_or_else(|| anyhow::anyhow!("service directory has no file name"))?;
let parent = std::fs::canonicalize(requested_parent)?;
let path = parent.join(file_name);
validate_private_service_directory(&parent)?;
match std::fs::symlink_metadata(&path) {
Ok(_) => {}
Err(error) if error.kind() == ErrorKind::NotFound => {
let mut builder = std::fs::DirBuilder::new();
builder.mode(mode).create(&path)?;
}
Err(error) => return Err(error.into()),
}
let mut options = OpenOptions::new();
options.read(true).custom_flags(
libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_DIRECTORY | libc::O_NONBLOCK,
);
let directory = options.open(&path)?;
// SAFETY: geteuid has no preconditions and does not retain pointers.
let effective_uid = unsafe { libc::geteuid() };
let metadata = directory.metadata()?;
if !metadata.is_dir() || metadata.uid() != effective_uid || metadata.mode() & 0o022 != 0 {
anyhow::bail!("service log directory has unsafe ownership or permissions");
}
directory.set_permissions(std::fs::Permissions::from_mode(mode))?;
directory.sync_all()?;
return Ok(());
}
#[cfg(not(unix))]
{
let _ = (path, mode);
anyhow::bail!("managed service directories require Unix filesystem checks")
}
}
fn open_private_service_log(path: &Path, mode: u32) -> anyhow::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt};
let requested_parent = path
.parent()
.ok_or_else(|| anyhow::anyhow!("service log path has no parent"))?;
let file_name = path
.file_name()
.ok_or_else(|| anyhow::anyhow!("service log path has no file name"))?;
let parent = std::fs::canonicalize(requested_parent)?;
let path = parent.join(file_name);
validate_private_service_directory(&parent)?;
let mut options = OpenOptions::new();
options
.create(true)
.append(true)
.mode(mode)
.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW | libc::O_NONBLOCK);
let file = options.open(&path)?;
// SAFETY: geteuid has no preconditions and does not retain pointers.
let effective_uid = unsafe { libc::geteuid() };
let metadata = file.metadata()?;
if !metadata.is_file() || metadata.uid() != effective_uid || metadata.nlink() != 1 {
anyhow::bail!("service log must be a regular, single-link file owned by root");
}
file.set_permissions(std::fs::Permissions::from_mode(mode))?;
file.sync_all()?;
return Ok(());
}
#[cfg(not(unix))]
{
let _ = (path, mode);
anyhow::bail!("managed service logs require Unix filesystem checks")
}
}
#[cfg(test)]
mod tests {
use super::{
ensure_private_service_directory, open_private_service_log, pick_bin, render_systemd_unit,
systemd_quote, validate_root_managed_service_file, validate_service_unit_path,
write_service_definition,
};
#[test]
fn systemd_unit_quotes_paths_without_changing_arguments() {
let unit = render_systemd_unit(
r#"/opt/Aether Tunnel/aether\"tunnel"#,
r#"/var/lib/aether tunnel/config\\node.toml"#,
"/var/lib/aether tunnel",
)
.expect("safe absolute paths should render");
assert!(unit.contains(r#"ExecStart="/opt/Aether Tunnel/aether\\\"tunnel""#));
assert!(unit.contains(
r#"Environment="AETHER_TUNNEL_CONFIG=/var/lib/aether tunnel/config\\\\node.toml""#
));
assert!(unit.contains(r#"WorkingDirectory="/var/lib/aether tunnel""#));
assert_eq!(systemd_quote("a\\b\"c"), r#""a\\b\"c""#);
}
#[test]
fn service_unit_paths_reject_directive_and_expansion_injection() {
for value in [
"relative/path",
"/tmp/config\nExecStart=/tmp/evil",
"/tmp/config\rEnvironment=EVIL=1",
"/tmp/%n/config",
"/tmp/$PATH/config",
] {
assert!(
validate_service_unit_path(value, "test").is_err(),
"accepted unsafe service path: {value:?}"
);
}
}
#[test]
fn service_commands_never_fall_back_to_path_lookup() {
assert_eq!(
pick_bin(&["relative-tool", "/definitely/missing/trusted-tool"]),
"/definitely/missing/trusted-tool"
);
}
#[cfg(unix)]
#[test]
fn root_service_rejects_files_beneath_shared_writable_directories() {
let directory =
std::env::temp_dir().join(format!("aether-service-path-test-{}", std::process::id()));
let _ = std::fs::remove_dir_all(&directory);
std::fs::create_dir(&directory).expect("test directory should be created");
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&directory, std::fs::Permissions::from_mode(0o777))
.expect("test directory permissions should be set");
}
let path = directory.join("aether-tunnel");
std::fs::write(&path, b"binary").expect("test binary should be written");
let result = validate_root_managed_service_file(&path, "binary", false);
let _ = std::fs::remove_dir_all(&directory);
assert!(result.is_err());
}
#[cfg(unix)]
#[test]
fn service_definitions_and_logs_refuse_links_and_use_private_atomic_files() {
use std::io::Read;
use std::os::unix::fs::{symlink, MetadataExt, PermissionsExt};
let directory = std::env::temp_dir().join(format!(
"aether-service-write-test-{}",
uuid::Uuid::new_v4()
));
std::fs::create_dir(&directory).unwrap();
std::fs::set_permissions(&directory, std::fs::Permissions::from_mode(0o700)).unwrap();
let definition = directory.join("aether-tunnel.service");
write_service_definition(definition.to_str().unwrap(), "first", 0o644).unwrap();
let metadata = std::fs::symlink_metadata(&definition).unwrap();
assert_eq!(metadata.mode() & 0o777, 0o644);
assert_eq!(metadata.nlink(), 1);
let mut old_definition = std::fs::File::open(&definition).unwrap();
write_service_definition(definition.to_str().unwrap(), "second", 0o644).unwrap();
let mut old_contents = String::new();
old_definition.read_to_string(&mut old_contents).unwrap();
assert_eq!(old_contents, "first");
assert_eq!(std::fs::read_to_string(&definition).unwrap(), "second");
let victim = directory.join("victim");
std::fs::write(&victim, b"known-good").unwrap();
std::fs::remove_file(&definition).unwrap();
symlink(&victim, &definition).unwrap();
assert!(write_service_definition(definition.to_str().unwrap(), "replace", 0o644).is_err());
assert_eq!(std::fs::read(&victim).unwrap(), b"known-good");
std::fs::remove_file(&definition).unwrap();
std::fs::hard_link(&victim, &definition).unwrap();
assert!(write_service_definition(definition.to_str().unwrap(), "replace", 0o644).is_err());
assert_eq!(std::fs::read(&victim).unwrap(), b"known-good");
std::fs::remove_file(&definition).unwrap();
let log_directory = directory.join("logs");
ensure_private_service_directory(&log_directory, 0o750).unwrap();
assert_eq!(
std::fs::symlink_metadata(&log_directory).unwrap().mode() & 0o777,
0o750
);
let log = log_directory.join("current.log");
open_private_service_log(&log, 0o640).unwrap();
let metadata = std::fs::symlink_metadata(&log).unwrap();
assert_eq!(metadata.mode() & 0o777, 0o640);
assert_eq!(metadata.nlink(), 1);
std::fs::remove_file(&log).unwrap();
symlink(&victim, &log).unwrap();
assert!(open_private_service_log(&log, 0o640).is_err());
assert_eq!(std::fs::read(&victim).unwrap(), b"known-good");
std::fs::remove_file(&log).unwrap();
std::fs::hard_link(&victim, &log).unwrap();
assert!(open_private_service_log(&log, 0o640).is_err());
assert_eq!(std::fs::read(&victim).unwrap(), b"known-good");
std::fs::remove_dir_all(directory).unwrap();
}
}
+2 -9
View File
@@ -200,11 +200,11 @@ impl App {
Field {
label: "Allow Private Targets",
key: "allow_private_targets",
value: "true".into(),
value: "false".into(),
kind: FieldKind::Bool,
required: false,
help:
"Allow proxying private/reserved upstream IPs by default; takes effect after restart",
"Allow proxying private/reserved upstream IPs; disabled by default and takes effect after restart",
},
Field {
label: "Heartbeat Interval",
@@ -450,13 +450,6 @@ impl App {
fn save(&mut self) -> anyhow::Result<()> {
let cfg = self.to_config()?;
cfg.save(&self.config_path)?;
// Restrict config file permissions to owner-only (contains management token).
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ =
std::fs::set_permissions(&self.config_path, std::fs::Permissions::from_mode(0o600));
}
self.modified = false;
self.saved_once = true;
self.message = Some((
File diff suppressed because it is too large Load Diff
+2
View File
@@ -55,6 +55,8 @@ pub struct ServerContext {
pub node_name: String,
/// Node ID assigned by this Aether server.
pub node_id: Arc<RwLock<String>>,
/// Server-issued node incarnation bound into tunnel authentication.
pub tunnel_generation: String,
/// API client for this server.
pub aether_client: Arc<AetherClient>,
/// Dynamic config from this server's heartbeat ACKs.
+169 -113
View File
@@ -1,5 +1,5 @@
use std::collections::{HashMap, HashSet};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant};
@@ -7,80 +7,11 @@ use tokio::sync::RwLock;
/// Check if an IP address belongs to a private/reserved network.
pub fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => is_private_ipv4(v4),
IpAddr::V6(v6) => is_private_ipv6(v6),
}
}
fn is_private_ipv4(ip: &Ipv4Addr) -> bool {
let octets = ip.octets();
// 10.0.0.0/8
if octets[0] == 10 {
return true;
}
// 172.16.0.0/12
if octets[0] == 172 && (16..=31).contains(&octets[1]) {
return true;
}
// 192.168.0.0/16
if octets[0] == 192 && octets[1] == 168 {
return true;
}
// 127.0.0.0/8
if octets[0] == 127 {
return true;
}
// 169.254.0.0/16 (link-local)
if octets[0] == 169 && octets[1] == 254 {
return true;
}
// 0.0.0.0/8
if octets[0] == 0 {
return true;
}
// 100.64.0.0/10 (CGNAT / shared address space)
if octets[0] == 100 && (64..=127).contains(&octets[1]) {
return true;
}
// 192.0.0.0/24 (IETF protocol assignments)
if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 {
return true;
}
// 198.18.0.0/15 (benchmark testing)
if octets[0] == 198 && (18..=19).contains(&octets[1]) {
return true;
}
// 240.0.0.0/4 (reserved for future use)
if octets[0] >= 240 {
return true;
}
false
}
fn is_private_ipv6(ip: &Ipv6Addr) -> bool {
// ::1 loopback
if ip.is_loopback() {
return true;
}
// :: unspecified
if ip.is_unspecified() {
return true;
}
let segments = ip.segments();
// fc00::/7 (ULA) - first byte is 0xfc or 0xfd
if segments[0] & 0xfe00 == 0xfc00 {
return true;
}
// fe80::/10 (link-local)
if segments[0] & 0xffc0 == 0xfe80 {
return true;
}
// IPv4-mapped IPv6 (::ffff:x.x.x.x) - check the embedded IPv4
if let Some(v4) = ip.to_ipv4_mapped() {
return is_private_ipv4(&v4);
}
false
// Keep every egress path on the same conservative classification as the
// gateway. This includes documentation, benchmarking, transition, and
// other reserved ranges that the standard `IpAddr::is_private` helpers do
// not cover.
aether_http::is_private_or_reserved_ip(*ip)
}
#[derive(Debug)]
@@ -108,6 +39,13 @@ impl std::fmt::Display for FilterError {
}
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
struct DnsCacheKey {
host: String,
port: u16,
allow_private: bool,
}
struct DnsCacheEntry {
addrs: Arc<Vec<SocketAddr>>,
expires_at: Instant,
@@ -115,12 +53,12 @@ struct DnsCacheEntry {
}
/// Lightweight DNS cache with TTL + capacity bounds.
/// Stores all public resolved addresses per host (used by SafeDnsResolver
/// to ensure reqwest connects to the same validated addresses).
/// Stores validated addresses per host and port so ACL checks and connection
/// setup can reuse the same resolution result.
pub struct DnsCache {
ttl: Duration,
capacity: usize,
entries: RwLock<HashMap<String, DnsCacheEntry>>,
entries: RwLock<HashMap<DnsCacheKey, DnsCacheEntry>>,
}
impl DnsCache {
@@ -132,31 +70,30 @@ impl DnsCache {
}
}
/// Look up cached public addresses for a host (any port).
/// Look up cached public addresses for a host + port.
///
/// Used by `SafeDnsResolver` which only knows the hostname — returns the
/// first unexpired entry whose key starts with `host:`.
pub async fn get_by_host(&self, host: &str) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
let prefix = format!("{}:", host.to_ascii_lowercase());
let now = Instant::now();
let entries = self.entries.read().await;
for (key, entry) in entries.iter() {
if key.starts_with(&prefix) && entry.expires_at > now {
return Some(Arc::clone(&entry.addrs));
}
}
None
/// This compatibility wrapper uses the restrictive (public-only) policy.
/// Callers that explicitly allow private targets must use
/// [`Self::get_for_policy`] so entries cannot cross the policy boundary.
#[allow(dead_code)]
pub async fn get(&self, host: &str, port: u16) -> Option<Arc<Vec<SocketAddr>>> {
self.get_for_policy(host, port, false).await
}
/// Look up cached public addresses for a host + port.
pub async fn get(&self, host: &str, port: u16) -> Option<Arc<Vec<SocketAddr>>> {
/// Look up a cached resolution under the exact target policy used to
/// validate it. Private-target and public-only resolutions are kept in
/// separate entries; otherwise a cache populated while private targets
/// are enabled could bypass filtering after a policy change.
pub async fn get_for_policy(
&self,
host: &str,
port: u16,
allow_private: bool,
) -> Option<Arc<Vec<SocketAddr>>> {
if self.capacity == 0 || self.ttl.is_zero() {
return None;
}
let key = Self::key(host, port);
let key = Self::key(host, port, allow_private);
let now = Instant::now();
// Fast path: read lock for cache hit
@@ -175,12 +112,26 @@ impl DnsCache {
None
}
/// Insert resolved public addresses into cache.
/// Insert resolved public addresses into the restrictive (public-only)
/// cache. This compatibility wrapper preserves the original API;
/// policy-aware callers should use [`Self::insert_for_policy`].
#[allow(dead_code)]
pub async fn insert(&self, host: &str, port: u16, addrs: Arc<Vec<SocketAddr>>) {
self.insert_for_policy(host, port, false, addrs).await;
}
/// Insert addresses under the exact target policy that produced them.
pub async fn insert_for_policy(
&self,
host: &str,
port: u16,
allow_private: bool,
addrs: Arc<Vec<SocketAddr>>,
) {
if self.capacity == 0 || self.ttl.is_zero() || addrs.is_empty() {
return;
}
let key = Self::key(host, port);
let key = Self::key(host, port, allow_private);
let now = Instant::now();
let mut entries = self.entries.write().await;
entries.retain(|_, entry| entry.expires_at > now);
@@ -205,8 +156,12 @@ impl DnsCache {
);
}
fn key(host: &str, port: u16) -> String {
format!("{}:{}", host.to_ascii_lowercase(), port)
fn key(host: &str, port: u16, allow_private: bool) -> DnsCacheKey {
DnsCacheKey {
host: host.to_ascii_lowercase(),
port,
allow_private,
}
}
}
@@ -222,16 +177,16 @@ pub async fn resolve_public_addrs(
dns_cache: &DnsCache,
) -> Result<Vec<SocketAddr>, FilterError> {
// Cache hit
if let Some(addrs) = dns_cache.get(host, port).await {
if let Some(addrs) = dns_cache.get_for_policy(host, port, allow_private).await {
return Ok((*addrs).clone());
}
// Async DNS resolution
let addr_str = format!("{}:{}", host, port);
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?
.collect();
// Async DNS resolution. Keep resolver wait time and answer count
// bounded before applying the private-address policy below.
let resolved: Vec<SocketAddr> =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?;
if resolved.is_empty() {
return Err(FilterError::DnsResolutionFailed(host.to_string()));
@@ -253,15 +208,17 @@ pub async fn resolve_public_addrs(
// Cache the validated public addresses
let arc_addrs = Arc::new(public);
dns_cache.insert(host, port, Arc::clone(&arc_addrs)).await;
dns_cache
.insert_for_policy(host, port, allow_private, Arc::clone(&arc_addrs))
.await;
Ok((*arc_addrs).clone())
}
/// Validate that the target host:port is allowed.
///
/// Performs port whitelist check, private IP filtering, and DNS resolution
/// with caching. The resolved addresses are stored in the shared DnsCache
/// so that the SafeDnsResolver can reuse them, eliminating the TOCTTOU gap.
/// with caching. The caller must use the returned addresses for the actual
/// connection rather than resolving the hostname again.
pub async fn validate_target(
host: &str,
port: u16,
@@ -282,12 +239,14 @@ pub async fn validate_target(
return Ok(vec![SocketAddr::new(ip, port)]);
}
// Resolve and validate DNS (populates cache for SafeDnsResolver)
// Resolve and return the exact addresses authorized for this request.
resolve_public_addrs(host, port, allow_private, dns_cache).await
}
#[cfg(test)]
mod tests {
use std::net::{Ipv4Addr, Ipv6Addr};
use super::*;
fn ports() -> HashSet<u16> {
@@ -318,9 +277,11 @@ mod tests {
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(198, 18, 0, 1))));
// Reserved
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(240, 0, 0, 1))));
// Multicast
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(224, 0, 0, 1))));
// Public
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1))));
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1))));
}
#[test]
@@ -335,6 +296,47 @@ mod tests {
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
0xfe80, 0, 0, 0, 0, 0, 0, 1
))));
// fec0::/10 (deprecated site-local)
assert!(is_private_ip(&"fec0::1".parse().unwrap()));
assert!(is_private_ip(
&"feff:ffff:ffff:ffff:ffff:ffff:ffff:ffff".parse().unwrap()
));
// ff00::/8 (multicast)
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
0xff02, 0, 0, 0, 0, 0, 0, 1
))));
// NAT64 well-known and local-use prefixes.
assert!(is_private_ip(&"64:ff9b::10.0.0.1".parse().unwrap()));
assert!(is_private_ip(&"64:ff9b::ffff:ffff".parse().unwrap()));
assert!(is_private_ip(&"64:ff9b:1::10.0.0.1".parse().unwrap()));
assert!(is_private_ip(
&"64:ff9b:1:ffff:ffff:ffff:ffff:ffff".parse().unwrap()
));
assert!(!is_private_ip(&"64:ff9a:ffff::1".parse().unwrap()));
assert!(!is_private_ip(&"64:ff9b:0:1::1".parse().unwrap()));
assert!(!is_private_ip(&"64:ff9b:2::1".parse().unwrap()));
// IPv6 transition formats with embedded IPv4 addresses.
assert!(is_private_ip(&"2002:0a00:0001::1".parse().unwrap()));
assert!(is_private_ip(
&"2002:ffff:ffff:ffff:ffff:ffff:ffff:ffff".parse().unwrap()
));
assert!(!is_private_ip(&"2003::1".parse().unwrap()));
assert!(is_private_ip(
&"2001:0000:4136:e378:8000:63bf:3fff:fdd2".parse().unwrap()
));
assert!(!is_private_ip(&"2001:1::1".parse().unwrap()));
assert!(is_private_ip(&"::192.0.2.1".parse().unwrap()));
assert!(is_private_ip(&"::ffff:0:192.0.2.1".parse().unwrap()));
assert!(is_private_ip(&"2001:db8::5efe:10.0.0.1".parse().unwrap()));
assert!(is_private_ip(
&"2001:db8::200:5efe:192.0.2.1".parse().unwrap()
));
// IPv4-mapped public addresses remain allowed, while private mapped
// addresses continue through the IPv4 classification.
assert!(!is_private_ip(&"::ffff:8.8.8.8".parse().unwrap()));
assert!(is_private_ip(&"::ffff:127.0.0.1".parse().unwrap()));
}
#[tokio::test]
@@ -351,6 +353,24 @@ mod tests {
assert!(matches!(result, Err(FilterError::PrivateIp(_))));
}
#[tokio::test]
async fn test_ipv6_site_local_blocked_unless_private_targets_allowed() {
let cache = cache();
let result = validate_target("fec0::1", 443, &ports(), false, &cache).await;
assert!(matches!(
result,
Err(FilterError::PrivateIp(IpAddr::V6(ip))) if ip == "fec0::1".parse::<Ipv6Addr>().unwrap()
));
let result = validate_target("fec0::1", 443, &ports(), true, &cache)
.await
.unwrap();
assert_eq!(
result,
vec![SocketAddr::new(IpAddr::V6("fec0::1".parse().unwrap()), 443)]
);
}
#[tokio::test]
async fn test_public_ip_allowed() {
let cache = cache();
@@ -412,4 +432,40 @@ mod tests {
let cached = cache.get("example.com", 443).await.unwrap();
assert_eq!(*cached, addrs);
}
#[tokio::test]
async fn test_cache_does_not_cross_private_target_policy() {
let cache = cache();
let private = Arc::new(vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 80)]);
let public = Arc::new(vec![SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)),
80,
)]);
cache
.insert_for_policy("example.com", 80, true, Arc::clone(&private))
.await;
assert!(cache
.get_for_policy("example.com", 80, false)
.await
.is_none());
assert_eq!(
*cache
.get_for_policy("example.com", 80, true)
.await
.expect("private-policy entry"),
*private
);
cache
.insert_for_policy("example.com", 80, false, Arc::clone(&public))
.await;
assert_eq!(
*cache
.get_for_policy("EXAMPLE.COM", 80, false)
.await
.expect("public-policy entry"),
*public
);
}
}
+136 -20
View File
@@ -3,7 +3,7 @@
use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use base64::Engine as _;
use tokio::net::TcpStream;
@@ -13,6 +13,7 @@ use tokio_tungstenite::tungstenite::http;
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use tracing::{debug, info, warn};
use crate::config::aether_url_for_log;
use crate::egress_proxy::{
connect_target_via_proxy, IpFamily, ProxyConnectOptions, UpstreamProxyConfig,
};
@@ -22,8 +23,10 @@ use aether_contracts::tunnel::{
TUNNEL_PROTOCOL_VERSION_HEADER,
};
use aether_contracts::tunnel_security::{
SecureFrameCodec, TunnelSecurityRole, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED,
TUNNEL_SECURITY_SESSION_HEADER,
sign_tunnel_security_handshake_for_generation, SecureFrameCodec, TunnelSecurityRole,
TUNNEL_GENERATION_HEADER, TUNNEL_SECURITY_HEADER, TUNNEL_SECURITY_NON_TLS_REQUIRED,
TUNNEL_SECURITY_PROOF_NONCE_HEADER, TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER,
TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER, TUNNEL_SECURITY_SESSION_HEADER,
};
use super::{dispatcher, heartbeat, writer};
@@ -48,7 +51,7 @@ pub async fn connect_and_run(
drain: watch::Receiver<bool>,
) -> Result<TunnelOutcome, anyhow::Error> {
let ws_url = build_tunnel_url(server);
debug!(url = %ws_url, conn = conn_idx, "connecting tunnel");
debug!(url = %aether_url_for_log(&ws_url), conn = conn_idx, "connecting tunnel");
// Build WebSocket request with auth headers
let mut request = ws_url.clone().into_client_request()?;
@@ -67,19 +70,38 @@ pub async fn connect_and_run(
);
let node_id = server.node_id.read().unwrap().clone();
insert_ascii_header(headers, "X-Node-Id", &node_id, "node_id")?;
let security_session = uuid::Uuid::new_v4().simple().to_string();
if server.tunnel_security == crate::config::TunnelSecurity::NonTlsRequired {
headers.insert(
TUNNEL_SECURITY_HEADER,
http::HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED),
);
insert_ascii_header(
headers,
TUNNEL_SECURITY_SESSION_HEADER,
&security_session,
"tunnel security session",
)?;
}
insert_ascii_header(
headers,
TUNNEL_GENERATION_HEADER,
&server.tunnel_generation,
"tunnel_generation",
)?;
let security_session =
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"))?;
let session = uuid::Uuid::new_v4().simple().to_string();
let nonce = uuid::Uuid::new_v4().simple().to_string();
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|_| anyhow::anyhow!("system clock is before the Unix epoch"))?
.as_secs();
insert_tunnel_security_handshake_headers(
headers,
key,
&node_id,
&server.tunnel_generation,
&session,
CURRENT_TUNNEL_PROTOCOL_VERSION,
timestamp,
&nonce,
)?;
session
} else {
String::new()
};
// 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.
@@ -449,9 +471,10 @@ async fn connect_direct_tunnel_tcp(
port: u16,
ip_family: IpFamily,
) -> io::Result<TcpStream> {
let resolved = tokio::net::lookup_host((host, port))
.await
.map_err(|err| io::Error::other(format!("tunnel DNS failed: {err}")))?;
let resolved =
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.map_err(|err| io::Error::other(format!("tunnel DNS failed: {err}")))?;
let addrs = filter_socket_addrs(resolved, ip_family);
if addrs.is_empty() {
@@ -540,6 +563,57 @@ fn insert_ascii_header(
Ok(())
}
fn insert_tunnel_security_handshake_headers(
headers: &mut http::HeaderMap,
key: &str,
node_id: &str,
tunnel_generation: &str,
session: &str,
protocol_version: u8,
timestamp_unix_secs: u64,
nonce: &str,
) -> anyhow::Result<()> {
let signature = sign_tunnel_security_handshake_for_generation(
key,
node_id,
tunnel_generation,
TUNNEL_SECURITY_NON_TLS_REQUIRED,
session,
protocol_version,
timestamp_unix_secs,
nonce,
)?;
headers.insert(
TUNNEL_SECURITY_HEADER,
http::HeaderValue::from_static(TUNNEL_SECURITY_NON_TLS_REQUIRED),
);
insert_ascii_header(
headers,
TUNNEL_SECURITY_SESSION_HEADER,
session,
"tunnel security session",
)?;
insert_ascii_header(
headers,
TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER,
&timestamp_unix_secs.to_string(),
"tunnel security proof timestamp",
)?;
insert_ascii_header(
headers,
TUNNEL_SECURITY_PROOF_NONCE_HEADER,
nonce,
"tunnel security proof nonce",
)?;
insert_ascii_header(
headers,
TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER,
&signature,
"tunnel security proof signature",
)?;
Ok(())
}
fn insert_node_name_headers(headers: &mut http::HeaderMap, node_name: &str) -> anyhow::Result<()> {
if node_name.is_ascii() {
return insert_ascii_header(headers, "X-Node-Name", node_name, "node_name");
@@ -620,4 +694,46 @@ mod tests {
.expect("encoded value should decode");
assert_eq!(decoded, "日本节点".as_bytes());
}
#[test]
fn tunnel_security_headers_include_verifiable_psk_proof() {
let key = base64::engine::general_purpose::STANDARD.encode([7_u8; 32]);
let mut headers = http::HeaderMap::new();
insert_tunnel_security_handshake_headers(
&mut headers,
&key,
"node-1",
"generation-1",
"0123456789abcdef0123456789abcdef",
CURRENT_TUNNEL_PROTOCOL_VERSION,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
)
.expect("security proof headers");
let signature = headers[TUNNEL_SECURITY_PROOF_SIGNATURE_HEADER]
.to_str()
.expect("signature header");
assert_eq!(
headers[TUNNEL_SECURITY_HEADER],
TUNNEL_SECURITY_NON_TLS_REQUIRED
);
assert_eq!(
headers[TUNNEL_SECURITY_PROOF_TIMESTAMP_HEADER],
"1700000000"
);
assert!(
aether_contracts::tunnel_security::verify_tunnel_security_handshake_for_generation(
&key,
"node-1",
"generation-1",
TUNNEL_SECURITY_NON_TLS_REQUIRED,
"0123456789abcdef0123456789abcdef",
CURRENT_TUNNEL_PROTOCOL_VERSION,
1_700_000_000,
"abcdef0123456789abcdef0123456789",
signature,
)
);
}
}
+223 -15
View File
@@ -1,13 +1,14 @@
//! Frame dispatcher: reads incoming WebSocket frames and routes them.
use std::collections::HashMap;
use std::collections::{HashMap, HashSet};
use std::mem::size_of;
use std::sync::Arc;
use std::sync::LazyLock;
use std::time::Duration;
use bytes::Bytes;
use futures_util::StreamExt;
use tokio::sync::mpsc;
use tokio::sync::watch;
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
use tokio::task::JoinHandle;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, info, warn};
@@ -15,12 +16,27 @@ use tracing::{debug, error, info, warn};
use crate::state::{AppState, ServerContext};
use super::heartbeat::HeartbeatHandle;
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
use super::protocol::{decompress_if_gzip_with_limit, Frame, MsgType, RequestMeta};
use super::stream_handler;
use super::stream_handler::StreamSendWindow;
use super::writer::FrameSender;
use aether_contracts::tunnel_security::SecureFrameCodec;
const REQUEST_BODY_QUEUE_BUDGET_BYTES: usize = 256 * 1024 * 1024;
static REQUEST_BODY_QUEUE_BUDGET: LazyLock<Arc<Semaphore>> =
LazyLock::new(|| Arc::new(Semaphore::new(REQUEST_BODY_QUEUE_BUDGET_BYTES)));
struct BudgetedFramePayload {
bytes: Bytes,
_permit: OwnedSemaphorePermit,
}
impl AsRef<[u8]> for BudgetedFramePayload {
fn as_ref(&self) -> &[u8] {
self.bytes.as_ref()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StreamDispatchStatus {
Delivered,
@@ -34,6 +50,24 @@ struct StreamDispatchTarget {
response_window: Arc<StreamSendWindow>,
}
/// A request stream is identified by a non-zero id and may only be opened
/// once while its handler is active. Replacing an entry in `streams` would
/// orphan the old body channel while still spawning another handler, making
/// the active-stream limit ineffective and allowing unbounded task growth.
fn validate_request_stream_id(
streams: &HashMap<u32, StreamDispatchTarget>,
active_handler_ids: &HashSet<u32>,
stream_id: u32,
) -> Result<(), &'static str> {
if stream_id == 0 {
return Err("invalid stream id");
}
if streams.contains_key(&stream_id) || active_handler_ids.contains(&stream_id) {
return Err("duplicate stream id");
}
Ok(())
}
/// Run the dispatcher loop, reading from the WebSocket stream.
#[allow(dead_code)]
pub async fn run<S>(
@@ -70,6 +104,10 @@ where
{
// Active streams: stream_id -> body sender + response flow-control window.
let mut streams: HashMap<u32, StreamDispatchTarget> = HashMap::new();
// A handler can outlive its routing entry when body dispatch fails. Keep
// its id reserved until the handler reports completion so a peer cannot
// reopen the same id and bypass the stream admission limit.
let mut active_handler_ids: HashSet<u32> = HashSet::new();
// Track spawned stream handlers so we can wait for them on shutdown
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
@@ -85,7 +123,7 @@ where
let mut draining = *drain.borrow();
let read_err = loop {
if draining && streams.is_empty() {
if draining && streams.is_empty() && active_handler_ids.is_empty() {
info!("tunnel drained after in-flight streams completed");
break None;
}
@@ -109,8 +147,9 @@ where
}
finished = handler_finished_rx.recv() => {
if let Some(stream_id) = finished {
active_handler_ids.remove(&stream_id);
streams.remove(&stream_id);
if draining && streams.is_empty() {
if draining && streams.is_empty() && active_handler_ids.is_empty() {
info!("tunnel drained after stream handler completion");
break None;
}
@@ -184,6 +223,20 @@ where
match frame.msg_type {
MsgType::RequestHeaders => {
if let Err(reason) =
validate_request_stream_id(&streams, &active_handler_ids, frame.stream_id)
{
warn!(
stream_id = frame.stream_id,
reason, "rejecting request headers with invalid stream id"
);
// Zero is reserved for connection-level control frames,
// so do not emit a stream-scoped error using that id.
if frame.stream_id != 0 {
try_send_stream_error(&frame_tx, frame.stream_id, reason);
}
continue;
}
if draining {
if frame_tx
.try_send(Frame::new(
@@ -203,7 +256,10 @@ where
}
// Decompress if the frame is gzip-compressed, then parse metadata
let payload = match decompress_if_gzip(&frame) {
let payload = match decompress_if_gzip_with_limit(
&frame,
aether_contracts::tunnel::MAX_TUNNEL_RELAY_META_LEN,
) {
Ok(p) => p,
Err(e) => {
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
@@ -233,7 +289,7 @@ where
}
};
if streams.len() >= max_streams {
if active_handler_ids.len() >= max_streams {
warn!(
stream_id = frame.stream_id,
"max concurrent streams reached"
@@ -267,6 +323,7 @@ where
response_window: Arc::clone(&response_window),
},
);
active_handler_ids.insert(frame.stream_id);
let request_headers_end_stream = frame.is_end_stream();
let state_clone = Arc::clone(&state);
@@ -321,7 +378,8 @@ where
"tunnel request body dispatch stalled",
);
}
if is_end && draining && streams.is_empty() {
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
{
info!("tunnel drained after request body completion");
break None;
}
@@ -333,7 +391,7 @@ where
// Client-side cancellation or end
if let Some(target) = streams.remove(&frame.stream_id) {
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
if draining && streams.is_empty() {
if draining && streams.is_empty() && active_handler_ids.is_empty() {
info!("tunnel drained after stream termination");
break None;
}
@@ -399,7 +457,7 @@ where
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
handler_handles.retain(|h| !h.is_finished());
frames_since_cleanup = 0;
if draining && streams.is_empty() {
if draining && streams.is_empty() && active_handler_ids.is_empty() {
info!("tunnel drained after cleanup");
break None;
}
@@ -421,12 +479,18 @@ where
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
let stream_id = frame.stream_id;
match tokio::time::timeout(stream_frame_dispatch_timeout(), tx.send(frame)).await {
Ok(Ok(())) => StreamDispatchStatus::Delivered,
Ok(Err(_)) => {
let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async {
let frame = attach_request_body_queue_budget(frame).await?;
tx.send(frame).await.ok()?;
Some(())
})
.await;
match dispatched {
Ok(Some(())) => StreamDispatchStatus::Delivered,
Ok(None) => {
warn!(
stream_id,
"stream handler channel closed while dispatching tunnel frame"
"stream handler channel or request body budget closed while dispatching tunnel frame"
);
StreamDispatchStatus::Closed
}
@@ -441,6 +505,50 @@ async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> Stream
}
}
async fn attach_request_body_queue_budget(frame: Frame) -> Option<Frame> {
attach_request_body_queue_budget_with(
frame,
Arc::clone(&REQUEST_BODY_QUEUE_BUDGET),
REQUEST_BODY_QUEUE_BUDGET_BYTES,
)
.await
}
async fn attach_request_body_queue_budget_with(
mut frame: Frame,
budget: Arc<Semaphore>,
budget_bytes: usize,
) -> Option<Frame> {
if frame.msg_type != MsgType::RequestBody {
return Some(frame);
}
let permits = request_body_queue_permits(&frame, budget_bytes)?;
let permit = budget.acquire_many_owned(permits).await.ok()?;
frame.payload = Bytes::from_owner(BudgetedFramePayload {
bytes: frame.payload,
_permit: permit,
});
Some(frame)
}
fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option<u32> {
let decoded_budget = if frame.is_gzip() {
aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES
} else {
0
};
let retained_bytes = frame
.payload
.len()
.checked_add(decoded_budget)?
.checked_add(size_of::<Frame>())?
.max(1);
if retained_bytes > budget_bytes {
return None;
}
u32::try_from(retained_bytes).ok()
}
/// Bound how long a single stream handler is allowed to block the shared
/// WebSocket read loop while receiving request-body frames.
fn stream_frame_dispatch_timeout() -> Duration {
@@ -497,6 +605,7 @@ async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
#[cfg(test)]
mod tests {
use super::*;
use aether_contracts::tunnel::{compress_payload, flags};
use aether_runtime::bounded_queue;
#[tokio::test]
@@ -534,6 +643,70 @@ mod tests {
assert_eq!(retained.payload, Bytes::from_static(b"first"));
}
#[tokio::test]
async fn request_body_queue_budget_releases_when_frame_is_dropped() {
const BUDGET_BYTES: usize = 4096;
let budget = Arc::new(Semaphore::new(BUDGET_BYTES));
let frame = Frame::new(
7,
MsgType::RequestBody,
0,
Bytes::from_static(b"request body"),
);
let permits = request_body_queue_permits(&frame, BUDGET_BYTES).expect("permit count");
let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), BUDGET_BYTES)
.await
.expect("frame should fit the queue budget");
assert_eq!(budget.available_permits(), BUDGET_BYTES - permits as usize);
drop(frame);
assert_eq!(budget.available_permits(), BUDGET_BYTES);
}
#[tokio::test]
async fn gzip_request_body_budget_follows_decoded_payload_lifetime() {
let (payload, frame_flags) = compress_payload(Bytes::from(vec![b'x'; 1024]));
assert_eq!(frame_flags, flags::GZIP_COMPRESSED);
let frame = Frame::new(7, MsgType::RequestBody, frame_flags, payload);
let required = request_body_queue_permits(&frame, REQUEST_BODY_QUEUE_BUDGET_BYTES)
.expect("gzip frame should fit the queue budget") as usize;
let budget = Arc::new(Semaphore::new(required));
let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), required)
.await
.expect("frame should acquire the entire local budget");
assert_eq!(budget.available_permits(), 0);
let decoded = stream_handler::decode_request_body_frame(frame)
.expect("gzip request body should decode");
assert_eq!(decoded, Bytes::from(vec![b'x'; 1024]));
assert_eq!(budget.available_permits(), 0);
drop(decoded);
assert_eq!(budget.available_permits(), required);
}
#[tokio::test]
async fn gzip_request_body_budget_releases_after_decode_error() {
let frame = Frame::new(
7,
MsgType::RequestBody,
flags::GZIP_COMPRESSED,
Bytes::from_static(b"not gzip"),
);
let required = request_body_queue_permits(&frame, REQUEST_BODY_QUEUE_BUDGET_BYTES)
.expect("gzip frame should fit the queue budget") as usize;
let budget = Arc::new(Semaphore::new(required));
let frame = attach_request_body_queue_budget_with(frame, Arc::clone(&budget), required)
.await
.expect("frame should acquire the entire local budget");
assert_eq!(budget.available_permits(), 0);
stream_handler::decode_request_body_frame(frame)
.expect_err("invalid gzip request body should fail");
assert_eq!(budget.available_permits(), required);
}
#[tokio::test]
async fn try_send_stream_error_emits_stream_error_frame() {
let (high_tx, mut high_rx) = bounded_queue::<Frame>(4);
@@ -581,4 +754,39 @@ mod tests {
assert!(!streams.contains_key(&7));
assert!(streams.contains_key(&9));
}
#[test]
fn request_stream_id_rejects_zero_and_active_duplicates() {
let (tx, _rx) = mpsc::channel::<Frame>(1);
let streams = HashMap::from([(
7,
StreamDispatchTarget {
body_tx: tx,
response_window: Arc::new(StreamSendWindow::new(1024)),
},
)]);
let mut active_handler_ids = HashSet::from([7]);
assert_eq!(
validate_request_stream_id(&streams, &active_handler_ids, 0),
Err("invalid stream id")
);
assert_eq!(
validate_request_stream_id(&streams, &active_handler_ids, 7),
Err("duplicate stream id")
);
assert_eq!(
validate_request_stream_id(&streams, &active_handler_ids, 9),
Ok(())
);
// The routing entry may be removed after a dispatch failure while the
// handler is still running; its reservation must continue to reject
// a new request with the same id.
active_handler_ids.insert(11);
assert_eq!(
validate_request_stream_id(&HashMap::new(), &active_handler_ids, 11),
Err("duplicate stream id")
);
}
}
+27 -3
View File
@@ -346,10 +346,12 @@ fn normalize_upgrade_target(raw: String) -> Option<String> {
.strip_prefix("tunnel-v")
.or_else(|| trimmed.strip_prefix("proxy-v"))
.unwrap_or(trimmed);
if normalized == CURRENT_VERSION {
let target = semver::Version::parse(normalized).ok()?;
let current = semver::Version::parse(CURRENT_VERSION).ok()?;
if target <= current {
return None;
}
Some(normalized.to_string())
Some(target.to_string())
}
fn maybe_trigger_upgrade(version: Option<String>) {
@@ -402,7 +404,10 @@ mod tests {
use arc_swap::ArcSwap;
use clap::Parser;
use super::{build_heartbeat_payload, handle_ack, AckDecision, HeartbeatSnapshot};
use super::{
build_heartbeat_payload, handle_ack, normalize_upgrade_target, AckDecision,
HeartbeatSnapshot, CURRENT_VERSION,
};
use crate::registration::client::AetherClient;
use crate::runtime::DynamicConfig;
use crate::state::{AppState, ServerContext, TunnelMetrics, TunnelRequestMetrics};
@@ -429,6 +434,7 @@ mod tests {
tunnel_encryption_key: config.tunnel_encryption_key.clone(),
node_name: config.node_name.clone(),
node_id: Arc::new(RwLock::new("node-123".to_string())),
tunnel_generation: "test-generation-1".to_string(),
aether_client: Arc::new(AetherClient::new(
&config,
&config.aether_url,
@@ -489,6 +495,24 @@ mod tests {
assert_eq!(server.dynamic.load().heartbeat_interval, 9);
}
#[test]
fn remote_upgrade_accepts_only_strict_semver_upgrades() {
let current = semver::Version::parse(CURRENT_VERSION).expect("package version is semver");
let target = semver::Version::new(current.major + 1, 0, 0);
assert_eq!(
normalize_upgrade_target(format!("tunnel-v{target}")),
Some(target.to_string())
);
assert_eq!(normalize_upgrade_target(CURRENT_VERSION.to_string()), None);
assert_eq!(normalize_upgrade_target("0.0.1".to_string()), None);
assert_eq!(
normalize_upgrade_target("1.2.3/../../payload".to_string()),
None
);
assert_eq!(normalize_upgrade_target("latest".to_string()), None);
}
#[tokio::test]
async fn heartbeat_payload_reports_resource_usage_and_tunnel_error_diagnostics() {
let config = sample_config();
+76 -4
View File
@@ -231,8 +231,14 @@ fn mix_u64(mut x: u64) -> u64 {
mod tests {
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Once};
use std::time::Duration;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_contracts::tunnel::{
sign_tunnel_relay_request, tunnel_relay_payload_digest, TUNNEL_RELAY_AUTH_NONCE_HEADER,
TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, TUNNEL_RELAY_AUTH_SENDER_HEADER,
TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER,
TUNNEL_RELAY_OWNER_INSTANCE_HEADER,
};
use aether_gateway::{build_router_with_state, AppState as GatewayAppState};
use arc_swap::ArcSwap;
use axum::Router;
@@ -303,7 +309,11 @@ mod tests {
.await
.expect("gateway should start");
let state = sample_state(sample_config(&gateway_base_url));
let mut tunnel_config = sample_config(&gateway_base_url);
tunnel_config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired;
tunnel_config.tunnel_encryption_key =
Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".to_string());
let state = sample_state(tunnel_config);
let server = sample_server(&state, "node-recovery");
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let tunnel_task = tokio::spawn({
@@ -370,12 +380,45 @@ mod tests {
gateway_base_url: &str,
node_id: &str,
) -> Option<(StatusCode, String)> {
let payload = relay_probe_envelope();
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("test clock should be after epoch")
.as_secs();
let nonce = uuid::Uuid::new_v4().simple().to_string();
let digest = tunnel_relay_payload_digest(&payload, &[]);
let signature = sign_tunnel_relay_request(
b"tunnel-reconnect-test-secret-at-least-32-bytes",
"tunnel-reconnect-test-client",
"tunnel-reconnect-test-gateway",
node_id,
"",
false,
timestamp,
&nonce,
&digest,
);
let response = reqwest::Client::new()
.post(format!(
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
))
.header("content-type", "application/octet-stream")
.body(relay_probe_envelope())
.header(
TUNNEL_RELAY_AUTH_SENDER_HEADER,
"tunnel-reconnect-test-client",
)
.header(
TUNNEL_RELAY_OWNER_INSTANCE_HEADER,
"tunnel-reconnect-test-gateway",
)
.header(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp)
.header(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce)
.header(
TUNNEL_RELAY_AUTH_PAYLOAD_HEADER,
digest.encode_header_value(),
)
.header(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature)
.body(payload)
.send()
.await
.ok()?;
@@ -411,12 +454,40 @@ mod tests {
async fn start_gateway_on_port(
port: u16,
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
let state = GatewayAppState::new().expect("gateway test state should build");
// The embedded gateway now fails closed when relay authentication is
// not configured. Keep this integration fixture explicitly authenticated.
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
std::env::set_var(
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
"tunnel-reconnect-test-secret-at-least-32-bytes",
);
std::env::set_var(
"AETHER_GATEWAY_INSTANCE_ID",
"tunnel-reconnect-test-gateway",
);
let mut state = GatewayAppState::new().expect("gateway test state should build");
aether_gateway::configure_test_tunnel_security(
&mut state,
"node-recovery",
"test-generation-1",
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
);
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
let router = build_router_with_state(state.clone());
let handle = spawn_router_on_port(port, router).await?;
Ok((state, handle))
}
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
if let Some(value) = value {
std::env::set_var(key, value);
} else {
std::env::remove_var(key);
}
}
async fn start_gateway_on_port_retry(
port: u16,
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
@@ -483,6 +554,7 @@ mod tests {
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())),
tunnel_generation: "test-generation-1".to_string(),
aether_client: Arc::new(AetherClient::new(
&config,
&config.aether_url,
File diff suppressed because it is too large Load Diff
+248 -134
View File
@@ -1,8 +1,7 @@
use std::collections::HashMap;
use std::convert::Infallible;
use std::future::Future;
use std::io;
use std::net::IpAddr;
use std::net::{IpAddr, SocketAddr};
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
@@ -17,7 +16,7 @@ use aether_contracts::{
use bytes::Bytes;
use futures_util::Stream;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::{BodyExt, Full, StreamBody};
use http_body_util::{BodyExt, StreamBody};
use hyper::body::Frame;
use hyper::rt;
use hyper::Response;
@@ -35,10 +34,9 @@ use tower_service::Service;
use crate::config::Config;
use crate::egress_proxy::{
connect_proxy_tcp, http_connect, socks5_connect, ProxyConnectOptions, UpstreamProxyConfig,
UpstreamProxyScheme,
connect_validated_target_via_proxy, ProxyConnectOptions, UpstreamProxyConfig,
};
use crate::target_filter::{self, DnsCache};
use crate::target_filter::DnsCache;
type BoxError = Box<dyn std::error::Error + Send + Sync>;
@@ -60,12 +58,75 @@ pub struct UpstreamClientPoolKey {
pub profile_id: String,
pub backend: String,
pub http_mode: String,
pub validated_target: ValidatedUpstreamTarget,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct ValidatedUpstreamTarget {
scheme: String,
host: String,
port: u16,
addrs: Vec<SocketAddr>,
}
impl ValidatedUpstreamTarget {
pub fn new(target_url: &url::Url, mut addrs: Vec<SocketAddr>) -> Result<Self, String> {
let scheme = target_url.scheme().to_ascii_lowercase();
if !matches!(scheme.as_str(), "http" | "https") {
return Err(format!("unsupported upstream scheme {scheme}"));
}
let host = target_url
.host_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| "missing host in upstream URL".to_string())?
.trim_start_matches('[')
.trim_end_matches(']')
.to_ascii_lowercase();
let port = target_url
.port_or_known_default()
.ok_or_else(|| "missing port in upstream URL".to_string())?;
if addrs.is_empty() {
return Err("validated upstream target has no addresses".to_string());
}
if addrs.iter().any(|addr| addr.port() != port) {
return Err("validated upstream target address has the wrong port".to_string());
}
addrs.sort_unstable();
addrs.dedup();
Ok(Self {
scheme,
host,
port,
addrs,
})
}
fn ensure_matches_uri(&self, uri: &Uri) -> Result<(), io::Error> {
let scheme = uri
.scheme_str()
.ok_or_else(|| io::Error::other("missing scheme"))?;
let host = uri_host(uri)?;
let port = uri_port_or_default(uri, scheme)?;
if !scheme.eq_ignore_ascii_case(&self.scheme)
|| !host.eq_ignore_ascii_case(&self.host)
|| port != self.port
{
return Err(io::Error::other(
"upstream connector target does not match its validated origin",
));
}
Ok(())
}
fn addrs(&self) -> &[SocketAddr] {
&self.addrs
}
}
#[derive(Clone)]
pub struct UpstreamClientPool {
config: Arc<Config>,
dns_cache: Arc<DnsCache>,
clients: Arc<Mutex<HashMap<UpstreamClientPoolKey, UpstreamClientPoolEntry>>>,
access_counter: Arc<AtomicU64>,
}
@@ -77,10 +138,9 @@ struct UpstreamClientPoolEntry {
}
impl UpstreamClientPool {
pub fn new(config: Arc<Config>, dns_cache: Arc<DnsCache>) -> Self {
pub fn new(config: Arc<Config>, _dns_cache: Arc<DnsCache>) -> Self {
Self {
config,
dns_cache,
clients: Arc::new(Mutex::new(HashMap::new())),
access_counter: Arc::new(AtomicU64::new(0)),
}
@@ -105,7 +165,7 @@ impl UpstreamClientPool {
.eq_ignore_ascii_case(TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE);
let client = build_upstream_client_with_protocol(
&self.config,
Arc::clone(&self.dns_cache),
key.validated_target.clone(),
http1_only,
h2c_prior_knowledge,
)?;
@@ -156,6 +216,7 @@ pub fn upstream_client_pool_key(
key_id: Option<&str>,
profile: Option<&ResolvedTransportProfile>,
http1_only: bool,
validated_target: ValidatedUpstreamTarget,
) -> UpstreamClientPoolKey {
let profile_http_mode = profile
.map(|profile| profile.http_mode.trim())
@@ -181,6 +242,7 @@ pub fn upstream_client_pool_key(
.unwrap_or(DEFAULT_BACKEND)
.to_string(),
http_mode: http_mode.to_string(),
validated_target,
}
}
@@ -201,18 +263,6 @@ fn validate_proxy_transport_backend(backend: &str) -> Result<(), String> {
Err(format!("unsupported transport profile backend: {backend}"))
}
pub fn http_proxy_authorization_header(proxy_url: Option<&str>) -> Option<String> {
let proxy = proxy_url
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| UpstreamProxyConfig::parse(value).ok())?;
if proxy.scheme() == UpstreamProxyScheme::Http {
proxy.basic_auth_header()
} else {
None
}
}
pub fn stream_request_body<S>(stream: S) -> UpstreamRequestBody
where
S: Stream<Item = Result<Frame<Bytes>, io::Error>> + Send + 'static,
@@ -220,9 +270,10 @@ where
StreamBody::new(stream).boxed_unsync()
}
#[cfg(test)]
pub fn full_request_body(body: Bytes) -> UpstreamRequestBody {
Full::new(body)
.map_err(|err: Infallible| match err {})
http_body_util::Full::new(body)
.map_err(|err: std::convert::Infallible| match err {})
.boxed_unsync()
}
@@ -241,18 +292,14 @@ pub struct RequestTiming {
pub connection_reused: bool,
}
#[derive(Clone)]
pub struct ValidatedResolver {
dns_cache: Arc<DnsCache>,
allow_private: bool,
#[derive(Clone, Debug)]
struct PinnedResolver {
target: ValidatedUpstreamTarget,
}
impl ValidatedResolver {
pub fn new(dns_cache: Arc<DnsCache>, allow_private: bool) -> Self {
Self {
dns_cache,
allow_private,
}
impl PinnedResolver {
fn new(target: ValidatedUpstreamTarget) -> Self {
Self { target }
}
}
@@ -268,7 +315,7 @@ impl Iterator for ValidatedAddrs {
}
}
impl Service<Name> for ValidatedResolver {
impl Service<Name> for PinnedResolver {
type Response = ValidatedAddrs;
type Error = io::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
@@ -278,22 +325,16 @@ impl Service<Name> for ValidatedResolver {
}
fn call(&mut self, name: Name) -> Self::Future {
let dns_cache = Arc::clone(&self.dns_cache);
let allow_private = self.allow_private;
let host = name.as_str().to_string();
let requested_host = name.as_str().to_string();
let target = self.target.clone();
Box::pin(async move {
if let Some(addrs) = dns_cache.get_by_host(&host).await {
return Ok(ValidatedAddrs {
inner: (*addrs).clone().into_iter(),
});
if !requested_host.eq_ignore_ascii_case(&target.host) {
return Err(io::Error::other(
"DNS request does not match the validated upstream host",
));
}
let resolved =
target_filter::resolve_public_addrs(&host, 0, allow_private, dns_cache.as_ref())
.await
.map_err(|err| io::Error::other(err.to_string()))?;
Ok(ValidatedAddrs {
inner: resolved.into_iter(),
inner: target.addrs.into_iter(),
})
})
}
@@ -301,9 +342,10 @@ impl Service<Name> for ValidatedResolver {
#[derive(Clone)]
pub struct InstrumentedConnector {
http: HttpConnector<ValidatedResolver>,
http: HttpConnector<PinnedResolver>,
tls_config: Arc<ClientConfig>,
proxy: Option<UpstreamProxyConfig>,
validated_target: ValidatedUpstreamTarget,
connect_timeout: Duration,
tcp_nodelay: bool,
tcp_keepalive: Option<Duration>,
@@ -319,9 +361,13 @@ impl Service<Uri> for InstrumentedConnector {
}
fn call(&mut self, dst: Uri) -> Self::Future {
if let Err(error) = self.validated_target.ensure_matches_uri(&dst) {
return Box::pin(async move { Err(error.into()) });
}
let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase());
let tls_config = Arc::clone(&self.tls_config);
if let Some(proxy) = self.proxy.clone() {
let validated_target = self.validated_target.clone();
let options = ProxyConnectOptions {
connect_timeout: self.connect_timeout,
tcp_nodelay: self.tcp_nodelay,
@@ -330,7 +376,16 @@ impl Service<Uri> for InstrumentedConnector {
};
let connect_start = std::time::Instant::now();
return Box::pin(async move {
connect_via_proxy(dst, scheme, tls_config, proxy, options, connect_start).await
connect_via_proxy(
dst,
scheme,
tls_config,
proxy,
validated_target,
options,
connect_start,
)
.await
});
}
let connecting = self.http.call(dst.clone());
@@ -381,39 +436,25 @@ async fn connect_via_proxy(
scheme: Option<String>,
tls_config: Arc<ClientConfig>,
proxy: UpstreamProxyConfig,
validated_target: ValidatedUpstreamTarget,
options: ProxyConnectOptions,
connect_start: std::time::Instant,
) -> Result<TimedConn, BoxError> {
let scheme = scheme.ok_or_else(|| io::Error::other("missing scheme"))?;
let target_host = uri_host(&dst)?;
let target_port = uri_port_or_default(&dst, &scheme)?;
let mut tcp = connect_proxy_tcp(
&proxy,
options.connect_timeout,
options.tcp_nodelay,
options.tcp_keepalive,
options.ip_family,
)
.await?;
match proxy.scheme() {
UpstreamProxyScheme::Http => {
if scheme == "https" {
http_connect(
&mut tcp,
&target_authority(&target_host, target_port),
&proxy,
)
.await?;
} else if scheme != "http" {
return Err(io::Error::other(format!("unsupported scheme {scheme}")).into());
let mut last_error = None;
let mut connected = None;
for target_addr in validated_target.addrs().iter().copied() {
match connect_validated_target_via_proxy(&proxy, target_addr, options).await {
Ok(tcp) => {
connected = Some(tcp);
break;
}
}
UpstreamProxyScheme::Socks5 | UpstreamProxyScheme::Socks5h => {
socks5_connect(&mut tcp, &proxy, &target_host, target_port).await?;
Err(error) => last_error = Some(error),
}
}
let tcp = connected.ok_or_else(|| {
last_error.unwrap_or_else(|| io::Error::other("validated upstream target has no addresses"))
})?;
let connect_ms = connect_start.elapsed().as_millis() as u64;
@@ -421,7 +462,7 @@ async fn connect_via_proxy(
"http" => Ok(TimedConn::new(
MaybeHttpsStream::Http {
stream: TokioIo::new(tcp),
is_proxy: proxy.scheme() == UpstreamProxyScheme::Http,
is_proxy: false,
},
ConnectTiming {
connect_ms,
@@ -465,24 +506,13 @@ fn uri_port_or_default(uri: &Uri, scheme: &str) -> Result<u16, io::Error> {
.ok_or_else(|| io::Error::other(format!("missing port for scheme {scheme}")))
}
fn target_authority(host: &str, port: u16) -> String {
if host.contains(':') && !host.starts_with('[') {
format!("[{host}]:{port}")
} else {
format!("{host}:{port}")
}
}
fn build_upstream_client_with_protocol(
config: &Config,
dns_cache: Arc<DnsCache>,
validated_target: ValidatedUpstreamTarget,
http1_only: bool,
h2c_prior_knowledge: bool,
) -> Result<UpstreamClient, String> {
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(
dns_cache,
config.allow_private_targets,
));
let mut http = HttpConnector::new_with_resolver(PinnedResolver::new(validated_target.clone()));
http.enforce_http(false);
http.set_connect_timeout(Some(Duration::from_secs(
config.upstream_connect_timeout_secs,
@@ -499,6 +529,7 @@ fn build_upstream_client_with_protocol(
let connector = InstrumentedConnector {
http,
tls_config: build_tls_config(http1_only),
validated_target,
proxy: config
.upstream_proxy_url
.as_deref()
@@ -791,12 +822,19 @@ mod tests {
header_fingerprint: None,
extra: None,
};
let target_url = url::Url::parse("https://example.com/").expect("target URL");
let validated_target = ValidatedUpstreamTarget::new(
&target_url,
vec![SocketAddr::from(([203, 0, 113, 10], 443))],
)
.expect("validated target");
let pool_key = upstream_client_pool_key(
Some("provider-1"),
Some("endpoint-1"),
Some("key-1"),
Some(&profile),
false,
validated_target,
);
assert_eq!(pool_key.provider_id, "provider-1");
@@ -852,25 +890,20 @@ mod tests {
assert!(!clients.contains_key(&key_b));
}
#[test]
fn http_proxy_authorization_header_uses_basic_auth_for_http_proxy() {
assert_eq!(
http_proxy_authorization_header(Some("http://user:[email protected]:8080")).as_deref(),
Some("Basic dXNlcjpwYXNz")
);
assert_eq!(
http_proxy_authorization_header(Some("socks5h://user:[email protected]:1080")),
None
);
}
fn test_pool_key(key_id: &str) -> UpstreamClientPoolKey {
let target_url = url::Url::parse("https://example.com/").expect("target URL");
let validated_target = ValidatedUpstreamTarget::new(
&target_url,
vec![SocketAddr::from(([203, 0, 113, 10], 443))],
)
.expect("validated target");
upstream_client_pool_key(
Some("provider-1"),
Some("endpoint-1"),
Some(key_id),
None,
false,
validated_target,
)
}
@@ -892,10 +925,10 @@ mod tests {
}
#[tokio::test]
#[ignore = "requires loopback listener support"]
async fn upstream_client_sends_http_requests_through_http_proxy() {
let (proxy_url, request_rx) = spawn_http_proxy().await;
let client = proxied_client(&proxy_url);
async fn http_proxy_connects_to_pinned_ip_and_preserves_origin_host() {
let pinned_addr = SocketAddr::from(([203, 0, 113, 77], 80));
let (proxy_url, connect_rx, request_rx) = spawn_http_proxy().await;
let client = proxied_client(&proxy_url, "http://example.com/", pinned_addr);
let request = hyper::Request::builder()
.method(hyper::Method::GET)
.uri("http://example.com/tunnel-test")
@@ -910,21 +943,66 @@ mod tests {
.await
.expect("body should collect")
.to_bytes();
let connect = connect_rx.await.expect("proxy should receive CONNECT");
let raw_request = request_rx.await.expect("proxy should receive request");
assert_eq!(status, hyper::StatusCode::OK);
assert_eq!(&body[..], b"ok");
assert!(
raw_request.starts_with("GET http://example.com/tunnel-test HTTP/1.1\r\n"),
connect.starts_with("CONNECT 203.0.113.77:80 HTTP/1.1\r\n"),
"unexpected proxy CONNECT: {connect:?}"
);
assert!(
raw_request.starts_with("GET /tunnel-test HTTP/1.1\r\n"),
"unexpected proxy request: {raw_request:?}"
);
assert!(
raw_request
.to_ascii_lowercase()
.contains("\r\nhost: example.com\r\n"),
"original Host header should be preserved: {raw_request:?}"
);
}
#[tokio::test]
#[ignore = "requires loopback listener support"]
async fn upstream_client_sends_http_requests_through_socks5h_proxy() {
let (proxy_url, target_rx, request_rx) = spawn_socks5h_proxy().await;
let client = proxied_client(&proxy_url);
async fn https_proxy_connects_to_pinned_ip_while_sni_uses_hostname() {
let pinned_addr = SocketAddr::from(([203, 0, 113, 78], 443));
let (proxy_url, connect_rx) = spawn_connect_only_http_proxy().await;
let client = proxied_client(&proxy_url, "https://sni.example/", pinned_addr);
let request = hyper::Request::builder()
.method(hyper::Method::GET)
.uri("https://sni.example/secure")
.body(full_request_body(Bytes::new()))
.expect("request should build");
let _ = client.request(request).await;
let connect = connect_rx.await.expect("proxy should receive CONNECT");
assert!(
connect.starts_with("CONNECT 203.0.113.78:443 HTTP/1.1\r\n"),
"unexpected proxy CONNECT: {connect:?}"
);
let uri: Uri = "https://sni.example/secure".parse().expect("URI");
match resolve_server_name(&uri).expect("server name") {
ServerName::DnsName(name) => assert_eq!(name.as_ref(), "sni.example"),
other => panic!("expected DNS SNI, got {other:?}"),
}
}
#[tokio::test]
async fn socks5_proxy_connects_to_pinned_ip_and_preserves_origin_host() {
assert_socks_proxy_uses_pinned_ip("socks5").await;
}
#[tokio::test]
async fn socks5h_proxy_connects_to_pinned_ip_and_preserves_origin_host() {
assert_socks_proxy_uses_pinned_ip("socks5h").await;
}
async fn assert_socks_proxy_uses_pinned_ip(scheme: &str) {
let pinned_addr = SocketAddr::from(([203, 0, 113, 79], 80));
let (proxy_addr, target_rx, request_rx) = spawn_socks5_proxy().await;
let proxy_url = format!("{scheme}://{proxy_addr}");
let client = proxied_client(&proxy_url, "http://example.com/", pinned_addr);
let request = hyper::Request::builder()
.method(hyper::Method::GET)
.uri("http://example.com/socks-test")
@@ -944,14 +1022,24 @@ mod tests {
.expect("SOCKS proxy should receive HTTP request");
assert_eq!(&body[..], b"ok");
assert_eq!(target, ("example.com".to_string(), 80));
assert_eq!(target, pinned_addr);
assert!(
raw_request.starts_with("GET /socks-test HTTP/1.1\r\n"),
"unexpected SOCKS tunneled request: {raw_request:?}"
);
assert!(
raw_request
.to_ascii_lowercase()
.contains("\r\nhost: example.com\r\n"),
"original Host header should be preserved: {raw_request:?}"
);
}
fn proxied_client(proxy_url: &str) -> UpstreamClient {
fn proxied_client(
proxy_url: &str,
target_url: &str,
pinned_addr: SocketAddr,
) -> UpstreamClient {
let _ = rustls::crypto::ring::default_provider().install_default();
let config = Config::try_parse_from([
"aether-tunnel",
@@ -967,23 +1055,32 @@ mod tests {
"2",
])
.expect("config should parse");
build_upstream_client_with_protocol(
&config,
Arc::new(DnsCache::new(Duration::from_secs(60), 16)),
true,
false,
)
.expect("client should build")
let target_url = url::Url::parse(target_url).expect("target URL should parse");
let validated_target = ValidatedUpstreamTarget::new(&target_url, vec![pinned_addr])
.expect("target should validate");
build_upstream_client_with_protocol(&config, validated_target, true, false)
.expect("client should build")
}
async fn spawn_http_proxy() -> (String, tokio::sync::oneshot::Receiver<String>) {
async fn spawn_http_proxy() -> (
String,
tokio::sync::oneshot::Receiver<String>,
tokio::sync::oneshot::Receiver<String>,
) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should exist");
let (connect_tx, connect_rx) = tokio::sync::oneshot::channel();
let (request_tx, request_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("proxy should accept");
let connect = read_http_headers(&mut stream).await;
let _ = connect_tx.send(connect);
stream
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await
.expect("CONNECT response should write");
let request = read_http_headers(&mut stream).await;
let _ = request_tx.send(request);
stream
@@ -991,12 +1088,30 @@ mod tests {
.await
.expect("proxy response should write");
});
(format!("http://{addr}"), request_rx)
(format!("http://{addr}"), connect_rx, request_rx)
}
async fn spawn_socks5h_proxy() -> (
String,
tokio::sync::oneshot::Receiver<(String, u16)>,
async fn spawn_connect_only_http_proxy() -> (String, tokio::sync::oneshot::Receiver<String>) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should exist");
let (connect_tx, connect_rx) = tokio::sync::oneshot::channel();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("proxy should accept");
let connect = read_http_headers(&mut stream).await;
let _ = connect_tx.send(connect);
stream
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await
.expect("CONNECT response should write");
});
(format!("http://{addr}"), connect_rx)
}
async fn spawn_socks5_proxy() -> (
SocketAddr,
tokio::sync::oneshot::Receiver<SocketAddr>,
tokio::sync::oneshot::Receiver<String>,
) {
let listener = TcpListener::bind("127.0.0.1:0")
@@ -1018,26 +1133,25 @@ mod tests {
.await
.expect("SOCKS method should write");
let mut request_head = [0u8; 5];
let mut request_head = [0u8; 4];
stream
.read_exact(&mut request_head)
.await
.expect("SOCKS request head should read");
assert_eq!(&request_head[..4], &[0x05, 0x01, 0x00, 0x03]);
let len = request_head[4] as usize;
let mut host = vec![0u8; len];
assert_eq!(request_head, [0x05, 0x01, 0x00, 0x01]);
let mut ip = [0u8; 4];
stream
.read_exact(&mut host)
.read_exact(&mut ip)
.await
.expect("SOCKS host should read");
.expect("SOCKS IPv4 target should read");
let mut port = [0u8; 2];
stream
.read_exact(&mut port)
.await
.expect("SOCKS port should read");
let host = String::from_utf8(host).expect("SOCKS host should be UTF-8");
let port = u16::from_be_bytes(port);
let _ = target_tx.send((host, port));
let target = SocketAddr::from((ip, port));
let _ = target_tx.send(target);
stream
.write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
.await
@@ -1050,7 +1164,7 @@ mod tests {
.await
.expect("SOCKS tunneled response should write");
});
(format!("socks5h://{addr}"), target_rx, request_rx)
(addr, target_rx, request_rx)
}
async fn read_http_headers(stream: &mut TcpStream) -> String {