diff --git a/Cargo.lock b/Cargo.lock index 8bbd5862..0ca21153 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -241,6 +241,16 @@ dependencies = [ "password-hash", ] +[[package]] +name = "ariadne" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36f5e3dca4e09a6f340a61a0e9c7b61e030c69fc27bf29d73218f7e5e3b7638f" +dependencies = [ + "unicode-width 0.1.11", + "yansi", +] + [[package]] name = "arrayvec" version = "0.7.6" @@ -915,7 +925,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bfcfdc083699101d5a7965e49925975f2f55060f94f9a05e7187be95d530ca59" dependencies = [ "once_cell", - "proc-macro-crate 3.2.0", + "proc-macro-crate 3.5.0", "proc-macro2", "quote", "syn 2.0.117", @@ -2234,6 +2244,7 @@ dependencies = [ "aes-gcm", "anyhow", "arc-swap", + "ariadne", "async-recursion", "async-ringbuf", "async-stream", @@ -2272,7 +2283,7 @@ dependencies = [ "gethostname 0.5.0", "git-version", "globwalk", - "guarden", + "guarden 0.2.0", "hickory-client", "hickory-proto", "hickory-resolver", @@ -2458,7 +2469,7 @@ dependencies = [ "dashmap", "easytier", "futures", - "guarden", + "guarden 0.1.2", "jsonwebtoken", "mimalloc", "mockall", @@ -3594,7 +3605,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ca87812d87fa82896df1adfb5c111cdeaae3edb6da028f5df002dcbd7df71454" dependencies = [ "futures", - "guarden-macros", + "guarden-macros 0.1.2", + "tokio", +] + +[[package]] +name = "guarden" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8408903291a7d0cc74169d5de4dd1919a9a402a2f67fcd7df3303ed045fae73" +dependencies = [ + "futures-core", + "guarden-macros 0.2.0", "tokio", ] @@ -3609,6 +3631,18 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "guarden-macros" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e0ef28f1077c259f9e7e238e234a78ce18cedbf0251fd2135f5fc23c40e79fe" +dependencies = [ + "proc-macro-crate 3.5.0", + "proc-macro2", + "quote", + "syn 2.0.117", +] + [[package]] name = "h2" version = "0.4.7" @@ -5579,7 +5613,7 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "680998035259dcfcafe653688bf2aa6d3e2dc05e98be6ab46afb089dc84f1df8" dependencies = [ - "proc-macro-crate 3.2.0", + "proc-macro-crate 2.0.0", "proc-macro2", "quote", "syn 2.0.117", @@ -6711,11 +6745,11 @@ dependencies = [ [[package]] name = "proc-macro-crate" -version = "3.2.0" +version = "3.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ecf48c7ca261d60b74ab1a7b20da18bede46776b2e55535cb958eb595c5fa7b" +checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" dependencies = [ - "toml_edit 0.22.20", + "toml_edit 0.25.12+spec-1.1.0", ] [[package]] @@ -7597,7 +7631,7 @@ checksum = "1f168d99749d307be9de54d23fd226628d99768225ef08f6ffb52e0182a27746" dependencies = [ "cfg-if", "glob", - "proc-macro-crate 3.2.0", + "proc-macro-crate 3.5.0", "proc-macro2", "quote", "regex", @@ -9921,8 +9955,7 @@ dependencies = [ [[package]] name = "tokio-websockets" version = "0.13.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dad543404f98bfc969aeb71994105c592acfc6c43323fddcd016bb208d1c65cb" +source = "git+https://github.com/EasyTier/tokio-websockets#dc9771c7c215882349c3cb328877550a3593df21" dependencies = [ "base64 0.22.1", "bytes", @@ -9997,6 +10030,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "toml_datetime" +version = "1.1.1+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7" +dependencies = [ + "serde_core", +] + [[package]] name = "toml_edit" version = "0.19.15" @@ -10034,6 +10076,18 @@ dependencies = [ "winnow 0.6.18", ] +[[package]] +name = "toml_edit" +version = "0.25.12+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2153edc6955a6c354fad8f5efd38b6a8769bdccf9fe50f8e1329f81b0baa5d7" +dependencies = [ + "indexmap 2.14.0", + "toml_datetime 1.1.1+spec-1.1.0", + "toml_parser", + "winnow 1.0.1", +] + [[package]] name = "toml_parser" version = "1.1.2+spec-1.1.0" @@ -11881,6 +11935,9 @@ name = "winnow" version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09dac053f1cd375980747450bfc7250c264eaae0583872e845c0c7cd578872b5" +dependencies = [ + "memchr", +] [[package]] name = "winreg" @@ -12262,7 +12319,7 @@ version = "5.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "897e79616e84aac4b2c46e9132a4f63b93105d54fe8c0e8f6bffc21fa8d49222" dependencies = [ - "proc-macro-crate 3.2.0", + "proc-macro-crate 3.5.0", "proc-macro2", "quote", "syn 2.0.117", @@ -12499,7 +12556,7 @@ version = "5.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5b59b012ebe9c46656f9cc08d8da8b4c726510aef12559da3e5f1bf72780752c" dependencies = [ - "proc-macro-crate 3.2.0", + "proc-macro-crate 3.5.0", "proc-macro2", "quote", "syn 2.0.117", diff --git a/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs b/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs index 26d898fb..69733bfe 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs +++ b/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs @@ -119,7 +119,9 @@ fn sync_tun_event_receivers(receivers: &mut HashMap) fn event_needs_tun_refresh(event: &GlobalCtxEvent) -> bool { matches!( event, - GlobalCtxEvent::DhcpIpv4Changed(_, _) | GlobalCtxEvent::ProxyCidrsUpdated(_, _) + GlobalCtxEvent::DhcpIpv4Changed(_, _) + | GlobalCtxEvent::ProxyCidrsUpdated(_, _) + | GlobalCtxEvent::PublicIpv6RoutesUpdated(_, _) ) } @@ -476,89 +478,68 @@ pub fn start_local_socket_server() -> bool { .retain(|instance_id, _| active_tun_candidate_ids.contains(instance_id)); shrink_hash_set_if_sparse(&mut delivered_tun_requests); shrink_hash_map_if_sparse(&mut last_tun_route_signatures); - let has_undelivered_tun_candidate = active_tun_candidate_ids - .iter() - .any(|instance_id| !delivered_tun_requests.contains(instance_id)); - let should_evaluate_tun = - tun_refresh || !tun_bootstrap_done || has_undelivered_tun_candidate; let mut saw_running_instance = false; let mut saw_tun_candidate = false; - if should_evaluate_tun { - for instance in snapshot.instances.iter() { - if instance.running { - saw_running_instance = true; - } - if instance.running && instance.tun_required { - saw_tun_candidate = true; - let virtual_ipv4 = instance - .my_node_info - .as_ref() - .and_then(|info| info.virtual_ipv4.clone()); - let virtual_ipv4_cidr = instance - .my_node_info - .as_ref() - .and_then(|info| info.virtual_ipv4_cidr.clone()); - if clients.is_empty() { - continue; - } - if virtual_ipv4.is_none() || virtual_ipv4_cidr.is_none() { - continue; - } - let aggregated_routes = aggregate_tun_routes(instance); - let route_signature = serde_json::to_string(&( - &virtual_ipv4, - &virtual_ipv4_cidr, - &aggregated_routes, - instance.magic_dns_enabled, - instance.need_exit_node, - )) - .unwrap_or_else(|_| "[]".to_string()); - let allow_route_signature_refresh = tun_refresh || !tun_bootstrap_done; - let should_send = !delivered_tun_requests.contains(&instance.instance_id) - || (allow_route_signature_refresh - && last_tun_route_signatures - .get(&instance.instance_id) - .map(|value| value != &route_signature) - .unwrap_or(true)); - if !should_send { - continue; - } - let payload = TunRequestPayload { - config_id: instance.config_id.clone(), - instance_id: instance.instance_id.clone(), - display_name: instance.display_name.clone(), - virtual_ipv4, - virtual_ipv4_cidr, - aggregated_routes, - magic_dns_enabled: instance.magic_dns_enabled, - need_exit_node: instance.need_exit_node, - }; - let payload_json = match serde_json::to_string(&payload) { - Ok(json) => json, - Err(err) => { - ohrs_log_error!("[Rust] serialize tun request failed: {}", err); - continue; - } - }; - if broadcast_local_socket_message( - &mut clients, - "tun_request", - &payload_json, - ) { - delivered_tun_requests.insert(instance.instance_id.clone()); - last_tun_route_signatures - .insert(instance.instance_id.clone(), route_signature); - } - } + for instance in snapshot.instances.iter() { + if instance.running { + saw_running_instance = true; } - } else { - for instance in snapshot.instances.iter() { - if instance.running { - saw_running_instance = true; - } - if instance.running && instance.tun_required { - saw_tun_candidate = true; + if !(instance.running && instance.tun_required) { + continue; + } + + saw_tun_candidate = true; + let virtual_ipv4 = instance + .my_node_info + .as_ref() + .and_then(|info| info.virtual_ipv4.clone()); + let virtual_ipv4_cidr = instance + .my_node_info + .as_ref() + .and_then(|info| info.virtual_ipv4_cidr.clone()); + if clients.is_empty() { + continue; + } + if virtual_ipv4.is_none() || virtual_ipv4_cidr.is_none() { + continue; + } + let aggregated_routes = aggregate_tun_routes(instance); + let route_signature = serde_json::to_string(&( + &virtual_ipv4, + &virtual_ipv4_cidr, + &aggregated_routes, + instance.magic_dns_enabled, + instance.need_exit_node, + )) + .unwrap_or_else(|_| "[]".to_string()); + let should_send = !delivered_tun_requests.contains(&instance.instance_id) + || last_tun_route_signatures + .get(&instance.instance_id) + .map(|value| value != &route_signature) + .unwrap_or(true); + if !should_send { + continue; + } + let payload = TunRequestPayload { + config_id: instance.config_id.clone(), + instance_id: instance.instance_id.clone(), + display_name: instance.display_name.clone(), + virtual_ipv4, + virtual_ipv4_cidr, + aggregated_routes, + magic_dns_enabled: instance.magic_dns_enabled, + need_exit_node: instance.need_exit_node, + }; + let payload_json = match serde_json::to_string(&payload) { + Ok(json) => json, + Err(err) => { + ohrs_log_error!("[Rust] serialize tun request failed: {}", err); + continue; } + }; + if broadcast_local_socket_message(&mut clients, "tun_request", &payload_json) { + delivered_tun_requests.insert(instance.instance_id.clone()); + last_tun_route_signatures.insert(instance.instance_id.clone(), route_signature); } } if !delivered_tun_requests.is_empty() diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index 2df2744b..19e560ce 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -51,7 +51,7 @@ time = "0.3" toml = "0.8.12" chrono = { version = "0.4.37", features = ["serde"] } -guarden = "0.1" +guarden = "0.2" delegate = "0.13.5" @@ -91,7 +91,7 @@ rustls = { version = "0.23.0", features = [ rcgen = { version = "0.12.1", optional = true } # for websocket -tokio-websockets = { version = "0.13.2", optional = true, features = [ +tokio-websockets = { version = "0.13.2", git = "https://github.com/EasyTier/tokio-websockets", optional = true, features = [ "rustls-webpki-roots", "client", "server", @@ -134,6 +134,7 @@ prost-wkt-types = "0.7.1" pbjson = "0.9.0" anyhow = "1.0" +ariadne = "0.5" url = { version = "2.5", features = ["serde"] } percent-encoding = "2.3.1" diff --git a/easytier/src/common/config.rs b/easytier/src/common/config.rs index b7183e38..9ddccc4f 100644 --- a/easytier/src/common/config.rs +++ b/easytier/src/common/config.rs @@ -6,6 +6,7 @@ use std::{ }; use anyhow::Context; +use ariadne::{CharSet, Config as AriadneConfig, IndexType, Label, Report, ReportKind, Source}; use base64::{Engine as _, prelude::BASE64_STANDARD}; use clap::ValueEnum; use clap::builder::PossibleValue; @@ -575,6 +576,35 @@ struct Config { source: Option, } +fn format_toml_parse_error(source_name: &str, config_str: &str, error: &toml::de::Error) -> String { + let message = format!("failed to parse config TOML from {source_name}"); + + let Some(span) = error.span() else { + return format!("{message}\ndetail: {error}"); + }; + + let mut output = Vec::new(); + let report = Report::build(ReportKind::Error, (source_name, span.clone())) + .with_config( + AriadneConfig::default() + .with_color(false) + .with_char_set(CharSet::Ascii) + .with_index_type(IndexType::Byte), + ) + .with_message(&message) + .with_label(Label::new((source_name, span)).with_message(error.message())) + .finish(); + + if report + .write((source_name, Source::from(config_str)), &mut output) + .is_ok() + { + String::from_utf8_lossy(&output).into_owned() + } else { + format!("{message}\ndetail: {error}") + } +} + #[derive(Debug, Clone)] pub struct TomlConfigLoader { config: Arc>, @@ -597,11 +627,35 @@ impl TomlConfigLoader { } pub fn new_from_str(config_str: &str) -> Result { - let mut config = toml::de::from_str::(config_str) - .with_context(|| format!("failed to parse config file: {}", config_str))?; + Self::new_from_str_with_source("inline config", config_str) + } + + pub fn new(config_path: &PathBuf) -> Result { + let config_str = std::fs::read_to_string(config_path) + .with_context(|| format!("failed to read config file: {}", config_path.display()))?; + + let source_name = config_path.display().to_string(); + Self::new_from_str_with_source(&source_name, &config_str) + } + + pub(crate) fn new_from_str_with_source( + source_name: &str, + config_str: &str, + ) -> Result { + let mut config = toml::de::from_str::(config_str).map_err(|err| { + let message = format_toml_parse_error(source_name, config_str, &err); + anyhow::Error::new(err).context(message) + })?; Self::normalize_config_source(&mut config); + Self::new_from_config(config).map_err(|err| { + let message = format!("failed to load config from {source_name}: {err}"); + err.context(message) + }) + } + + fn new_from_config(mut config: Config) -> Result { config.flags_struct = Some( Self::gen_flags(config.flags.clone().unwrap_or_default()) .context("failed to parse flags")?, @@ -637,14 +691,6 @@ impl TomlConfigLoader { Ok(config) } - pub fn new(config_path: &PathBuf) -> Result { - let config_str = std::fs::read_to_string(config_path) - .with_context(|| format!("failed to read config file: {:?}", config_path))?; - let ret = Self::new_from_str(&config_str)?; - - Ok(ret) - } - fn gen_flags( flags_hashmap: serde_json::Map, ) -> serde_json::Result { @@ -1221,13 +1267,13 @@ pub async fn load_config_from_file( .read_to_string(&mut stdin) .await .context("failed to read config from stdin")?; - let config = TomlConfigLoader::new_from_str(&stdin)?; + let config = TomlConfigLoader::new_from_str_with_source("stdin", &stdin)?; return Ok((config, ConfigFileControl::STATIC_CONFIG)); } let config_str = tokio::fs::read_to_string(config_file) .await - .with_context(|| format!("failed to read config file: {:?}", config_file))?; + .with_context(|| format!("failed to read config file: {}", config_file.display()))?; let (expanded_config_str, uses_env_vars) = if disable_env_parsing { (config_str.clone(), false) @@ -1249,8 +1295,8 @@ pub async fn load_config_from_file( ); } - let config = TomlConfigLoader::new_from_str(&expanded_config_str) - .with_context(|| format!("failed to load config file: {:?}", config_file))?; + let source_name = config_file.display().to_string(); + let config = TomlConfigLoader::new_from_str_with_source(&source_name, &expanded_config_str)?; let mut control = ConfigFileControl::from_path(config_file.clone()).await; @@ -1290,6 +1336,96 @@ pub mod tests { use std::path::PathBuf; use tempfile::NamedTempFile; + #[test] + fn invalid_toml_error_includes_location_and_source_line() { + let error = TomlConfigLoader::new_from_str("dhcp = \"yes\"").unwrap_err(); + let display = error.to_string(); + + assert!(display.contains("failed to parse config TOML")); + assert!(display.contains("inline config")); + assert!(display.contains("dhcp = \"yes\"")); + assert!(display.contains("^")); + assert!(display.contains("invalid type: string")); + assert!(!display.contains("")); + assert!( + error + .chain() + .any(|err| err.downcast_ref::().is_some()) + ); + } + + #[test] + fn invalid_file_toml_error_includes_config_source() { + let mut config_file = NamedTempFile::new().unwrap(); + writeln!(config_file, "dhcp = \"yes\"").unwrap(); + + let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err(); + let error = error.to_string(); + + assert!(error.contains(config_file.path().to_string_lossy().as_ref())); + assert!(error.contains("failed to parse config TOML")); + assert!(error.contains("dhcp = \"yes\"")); + assert!(error.contains("^")); + assert!(error.contains("invalid type: string")); + assert!(!error.contains("")); + } + + #[test] + fn invalid_stdin_toml_error_includes_config_source_in_display() { + let error = TomlConfigLoader::new_from_str_with_source("stdin", "dhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("stdin")); + assert!(error.contains("failed to parse config TOML")); + assert!(error.contains("dhcp = \"yes\"")); + assert!(error.contains("^")); + assert!(error.contains("invalid type: string")); + assert!(!error.contains("")); + } + + #[test] + fn invalid_toml_error_handles_non_ascii_before_error() { + let error = TomlConfigLoader::new_from_str("hostname = \"节点\"\ndhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("dhcp = \"yes\"")); + assert!(error.contains("^")); + assert!(error.contains("invalid type: string")); + } + + #[test] + fn invalid_toml_error_handles_non_ascii_before_error_on_same_line() { + let error = TomlConfigLoader::new_from_str("hostname = \"节点\" dhcp = \"yes\"") + .unwrap_err() + .to_string(); + + assert!(error.contains("failed to parse config TOML")); + assert!(error.contains("inline config:1:")); + assert!(error.contains("hostname = \"节点\" dhcp = \"yes\"")); + assert!(error.contains("expected newline")); + assert!(!error.contains("")); + } + + #[test] + fn invalid_file_flags_error_includes_config_source_in_display() { + let mut config_file = NamedTempFile::new().unwrap(); + writeln!(config_file, "[flags]").unwrap(); + writeln!(config_file, "socket_mark = \"bad\"").unwrap(); + + let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err(); + + let display = error.to_string(); + assert!(display.contains(config_file.path().to_string_lossy().as_ref())); + assert!(display.contains("failed to load config")); + assert!(display.contains("failed to parse flags")); + + // with_context preserves the cause chain so callers can inspect the root reason. + let chain: Vec = error.chain().map(|e| e.to_string()).collect(); + assert!(chain.iter().any(|m| m.contains("failed to parse flags"))); + } + #[test] fn socket_mark_config_file_roundtrip_none_some_and_zero() { // Omitting the flag leaves socket_mark unset (None) -> SO_MARK untouched. diff --git a/easytier/src/core.rs b/easytier/src/core.rs index 5c4d7aa2..bee76cae 100644 --- a/easytier/src/core.rs +++ b/easytier/src/core.rs @@ -1644,7 +1644,7 @@ pub async fn main() -> ExitCode { // Verify configurations if cli.check_config { if let Err(error) = validate_config(&cli).await { - log::error!(?error, "Config validation failed"); + log::error!(%error, "Config validation failed"); return ExitCode::FAILURE; } else { return ExitCode::SUCCESS; @@ -1654,7 +1654,7 @@ pub async fn main() -> ExitCode { let mut ret_code = 0; if let Err(error) = run_main(cli).await { - log::error!(?error); + log::error!(%error); ret_code = 1; } @@ -1674,12 +1674,13 @@ async fn validate_config(cli: &Cli) -> anyhow::Result<()> { for config_file in config_files { if config_file == &PathBuf::from("-") { let mut stdin = String::new(); - _ = tokio::io::stdin().read_to_string(&mut stdin).await?; - TomlConfigLoader::new_from_str(stdin.as_str()) - .with_context(|| "config source: stdin")?; + _ = tokio::io::stdin() + .read_to_string(&mut stdin) + .await + .context("failed to read config from stdin")?; + TomlConfigLoader::new_from_str_with_source("stdin", stdin.as_str())?; } else { - TomlConfigLoader::new(config_file) - .with_context(|| format!("config source: {:?}", config_file))?; + TomlConfigLoader::new(config_file)?; }; } diff --git a/easytier/src/peers/foreign_network_manager.rs b/easytier/src/peers/foreign_network_manager.rs index c34240cc..9a9fb641 100644 --- a/easytier/src/peers/foreign_network_manager.rs +++ b/easytier/src/peers/foreign_network_manager.rs @@ -14,7 +14,7 @@ use std::{ }; use dashmap::{DashMap, DashSet}; -use guarden::defer; +use guarden::{Guard, defer}; use tokio::{ sync::{ Mutex, diff --git a/easytier/src/peers/peer_conn.rs b/easytier/src/peers/peer_conn.rs index d7fdbde8..90374286 100644 --- a/easytier/src/peers/peer_conn.rs +++ b/easytier/src/peers/peer_conn.rs @@ -33,7 +33,6 @@ use super::{ peer_session::{PeerSession, PeerSessionAction}, traffic_metrics::AggregateTrafficMetrics, }; -use crate::utils::BoxExt; use crate::{ common::{ PeerId, @@ -380,9 +379,9 @@ impl PeerConn { session_filter, noise_handshake_result: None, - tunnel: Arc::new(Mutex::new( - guard!([mut mpsc_tunnel] mpsc_tunnel.close()).boxed(), - )), + tunnel: Arc::new(Mutex::new(Box::new( + guard!([mut mpsc_tunnel] mpsc_tunnel.close()), + ))), sink, recv: Mutex::new(Some(recv)), tunnel_info, diff --git a/easytier/src/peers/peer_manager.rs b/easytier/src/peers/peer_manager.rs index 2a76e654..17144d4c 100644 --- a/easytier/src/peers/peer_manager.rs +++ b/easytier/src/peers/peer_manager.rs @@ -1533,9 +1533,22 @@ impl PeerManager { ) -> Result<(), Error> { let policy = Self::get_next_hop_policy(msg.peer_manager_header().unwrap().is_latency_first()); + let is_latency_first = msg.peer_manager_header().unwrap().is_latency_first(); let packet_type = msg.peer_manager_header().unwrap().packet_type; let msg_len = msg.buf_len() as u64; - let send_result = if peers.has_peer(dst_peer_id) { + let latency_first_gateway = if is_latency_first { + peers + .get_gateway_peer_id(dst_peer_id, policy.clone()) + .await + .filter(|gateway| *gateway != dst_peer_id) + } else { + None + }; + let send_result = if let Some(gateway) = latency_first_gateway + && (peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway)) + { + relay_peer_map.send_msg(msg, dst_peer_id, policy).await + } else if peers.has_peer(dst_peer_id) { peers.send_msg_directly(msg, dst_peer_id).await } else if foreign_network_client.has_next_hop(dst_peer_id) { foreign_network_client.send_msg(msg, dst_peer_id).await @@ -2185,6 +2198,7 @@ impl PeerManager { mod tests { use base64::Engine; use std::{ + collections::HashMap, fmt::Debug, sync::Arc, time::{Duration, Instant}, @@ -2192,6 +2206,7 @@ mod tests { use crate::{ common::{ + PeerId, config::Flags, global_ctx::{NetworkIdentity, tests::get_mock_global_ctx}, stats_manager::{LabelSet, LabelType, MetricName}, @@ -2206,7 +2221,7 @@ mod tests { peer_conn::tests::set_secure_mode_cfg, peer_manager::RouteAlgoType, peer_rpc::tests::register_service, - route_trait::NextHopPolicy, + route_trait::{NextHopPolicy, RouteCostCalculatorInterface}, tests::{ connect_peer_manager, create_mock_peer_manager_with_name, wait_route_appear, wait_route_appear_with_cost, @@ -2250,6 +2265,16 @@ mod tests { )) } + struct TestCostCalculator { + costs: HashMap<(PeerId, PeerId), i32>, + } + + impl RouteCostCalculatorInterface for TestCostCalculator { + fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 { + *self.costs.get(&(src, dst)).unwrap_or(&1) + } + } + #[test] fn recent_traffic_fanout_policy_only_marks_single_peer() { assert!(PeerManager::should_mark_recent_traffic_for_fanout(0)); @@ -2657,6 +2682,109 @@ mod tests { .await; } + #[tokio::test] + async fn send_msg_internal_uses_latency_first_gateway_for_direct_peer() { + let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; + let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; + let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; + + connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await; + connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await; + connect_peer_manager(peer_mgr_a.clone(), peer_mgr_c.clone()).await; + wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone()) + .await + .unwrap(); + wait_route_appear(peer_mgr_b.clone(), peer_mgr_c.clone()) + .await + .unwrap(); + wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone()) + .await + .unwrap(); + + peer_mgr_a + .get_route() + .set_route_cost_fn(Box::new(TestCostCalculator { + costs: HashMap::from([ + ((peer_mgr_a.my_peer_id(), peer_mgr_c.my_peer_id()), 100), + ((peer_mgr_a.my_peer_id(), peer_mgr_b.my_peer_id()), 1), + ((peer_mgr_b.my_peer_id(), peer_mgr_c.my_peer_id()), 1), + ]), + })) + .await; + + wait_for_condition( + || { + let peer_mgr_a = peer_mgr_a.clone(); + let peer_mgr_b = peer_mgr_b.clone(); + let peer_mgr_c = peer_mgr_c.clone(); + async move { + peer_mgr_a + .get_route() + .get_next_hop_with_policy(peer_mgr_c.my_peer_id(), NextHopPolicy::LeastCost) + .await + == Some(peer_mgr_b.my_peer_id()) + } + }, + Duration::from_secs(5), + ) + .await; + + let b_network_labels = network_labels(&peer_mgr_b); + let forwarded_bytes_before = metric_value( + &peer_mgr_b, + MetricName::TrafficBytesForwarded, + &b_network_labels, + ); + let forwarded_packets_before = metric_value( + &peer_mgr_b, + MetricName::TrafficPacketsForwarded, + &b_network_labels, + ); + + let mut pkt = ZCPacket::new_with_payload(b"latency-first"); + pkt.fill_peer_manager_hdr( + peer_mgr_a.my_peer_id(), + peer_mgr_c.my_peer_id(), + PacketType::Data as u8, + ); + pkt.mut_peer_manager_header() + .unwrap() + .set_latency_first(true); + let pkt_len = pkt.buf_len() as u64; + + PeerManager::send_msg_internal( + &peer_mgr_a.peers, + &peer_mgr_a.foreign_network_client, + &peer_mgr_a.relay_peer_map, + Some(&peer_mgr_a.traffic_metrics), + pkt, + peer_mgr_c.my_peer_id(), + ) + .await + .unwrap(); + + wait_for_condition( + || { + let peer_mgr_b = peer_mgr_b.clone(); + let b_network_labels = b_network_labels.clone(); + async move { + metric_value( + &peer_mgr_b, + MetricName::TrafficBytesForwarded, + &b_network_labels, + ) >= forwarded_bytes_before + pkt_len + && metric_value( + &peer_mgr_b, + MetricName::TrafficPacketsForwarded, + &b_network_labels, + ) > forwarded_packets_before + } + }, + Duration::from_secs(5), + ) + .await; + } + #[tokio::test] async fn send_msg_internal_records_control_metrics_for_direct_peer() { let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await; diff --git a/easytier/src/tunnel/fake_tcp/mod.rs b/easytier/src/tunnel/fake_tcp/mod.rs index 0f27d84f..9cd052f8 100644 --- a/easytier/src/tunnel/fake_tcp/mod.rs +++ b/easytier/src/tunnel/fake_tcp/mod.rs @@ -12,7 +12,7 @@ use std::{ sync::Arc, task::{Context as TaskContext, Poll}, }; -use tokio::{io::AsyncReadExt, net::TcpStream, sync::Mutex}; +use tokio::{io::AsyncReadExt, net::TcpStream}; use crate::tunnel::{ FromUrl, IpVersion, SinkError, SinkItem, StreamItem, Tunnel, TunnelConnector, TunnelError, @@ -85,7 +85,7 @@ pub struct FakeTcpTunnelListener { addr: url::Url, os_listener: Option, // interface_name -> fake tcp stack - stack_map: DashMap>>, + stack_map: DashMap>, // a cache from ip addr to interface name ip_to_ifname: IpToIfNameCache, } @@ -148,7 +148,7 @@ impl FakeTcpTunnelListener { async fn get_stack( &self, accept_result: &AcceptResult, - ) -> Result>, TunnelError> { + ) -> Result, TunnelError> { let local_socket_addr = accept_result.local_addr; let interface_name = &accept_result.interface_name; @@ -158,29 +158,38 @@ impl FakeTcpTunnelListener { IpAddr::V6(ip) => (None, Some(ip)), }; - let ret = match self.stack_map.entry(interface_name.to_string()) { - dashmap::Entry::Occupied(entry) => entry.get().clone(), - dashmap::Entry::Vacant(entry) => { - let tun = - create_tun_off_runtime(interface_name.to_string(), None, local_socket_addr) - .await?; - tracing::info!( - ?local_socket_addr, - "create new stack with interface_name: {:?}", - interface_name - ); - let stack = Arc::new(Mutex::new(stack::Stack::new( - tun, - local_ip.unwrap_or(Ipv4Addr::UNSPECIFIED), - local_ip6, - accept_result.mac, - ))); - entry.insert(stack.clone()); - stack - } - }; + if let Some(entry) = self.stack_map.get(interface_name) { + let stack = entry.clone(); + drop(entry); - Ok(ret) + if !stack.is_closed() { + return Ok(stack); + } + + tracing::warn!( + interface_name, + "fake_tcp stack reader_task finished, recreating stack" + ); + self.stack_map.remove(interface_name); + } + + let tun = + create_tun_off_runtime(interface_name.to_string(), None, local_socket_addr).await?; + tracing::info!( + ?local_socket_addr, + "create new stack with interface_name: {:?}", + interface_name + ); + let stack = Arc::new(stack::Stack::new( + tun, + local_ip.unwrap_or(Ipv4Addr::UNSPECIFIED), + local_ip6, + accept_result.mac, + )); + self.stack_map + .insert(interface_name.to_string(), stack.clone()); + + Ok(stack) } } @@ -215,19 +224,29 @@ impl TunnelListener for FakeTcpTunnelListener { let os_listener = tokio::net::TcpListener::bind(bind_addr).await?; tracing::info!(port, "FakeTcpTunnelListener listening"); self.os_listener = Some(os_listener); - // self.stack.lock().await.listen(port); Ok(()) } async fn accept(&mut self) -> Result, TunnelError> { tracing::debug!("FakeTcpTunnelListener waiting for accept"); - let res = self.do_accept().await?; - let stack = self.get_stack(&res).await?; - let socket = stack - .lock() - .await - .alloc_established_socket(res.local_addr, res.remote_addr, stack::State::Established) - .await; + let (res, stack, socket) = loop { + let res = self.do_accept().await?; + let stack = self.get_stack(&res).await?; + let socket = stack.try_alloc_established_socket( + res.local_addr, + res.remote_addr, + stack::State::Established, + ); + let Some(socket) = socket else { + tracing::warn!( + interface_name = res.interface_name, + "fake_tcp stack closed while accepting connection, dropping accepted socket" + ); + self.stack_map.remove(&res.interface_name); + continue; + }; + break (res, stack, socket); + }; tracing::info!( ?res, @@ -236,7 +255,7 @@ impl TunnelListener for FakeTcpTunnelListener { ); let info = TunnelInfo { - tunnel_type: get_faketcp_tunnel_type_str(stack.lock().await.driver_type()), + tunnel_type: get_faketcp_tunnel_type_str(stack.driver_type()), local_addr: Some(self.local_url().into()), remote_addr: Some( crate::tunnel::build_url_from_socket_addr( @@ -354,12 +373,14 @@ impl TunnelConnector for FakeTcpTunnelConnector { let tun = create_tun_off_runtime(interface_name.clone(), Some(remote_addr), local_addr).await?; let local_ip = local_ip.unwrap_or("0.0.0.0".parse().unwrap()); - let mut stack = stack::Stack::new(tun, local_ip, local_ip6, mac); + let stack = stack::Stack::new(tun, local_ip, local_ip6, mac); let driver_type = stack.driver_type(); let socket = stack - .alloc_established_socket(local_addr, remote_addr, stack::State::SynSent) - .await; + .try_alloc_established_socket(local_addr, remote_addr, stack::State::SynSent) + .ok_or(TunnelError::InternalError( + "FakeTCP stack closed while allocating socket".into(), + ))?; let os_stream = os_socket.connect(remote_addr).await?; diff --git a/easytier/src/tunnel/fake_tcp/stack.rs b/easytier/src/tunnel/fake_tcp/stack.rs index 6866c926..a7f1b779 100644 --- a/easytier/src/tunnel/fake_tcp/stack.rs +++ b/easytier/src/tunnel/fake_tcp/stack.rs @@ -54,7 +54,7 @@ use std::sync::{ use tokio::sync::broadcast; use tokio::time; use tokio_util::task::AbortOnDropHandle; -use tracing::{info, trace, warn}; +use tracing::{error, info, trace, warn}; const TIMEOUT: time::Duration = time::Duration::from_secs(1); const RETRIES: usize = 6; @@ -83,13 +83,33 @@ impl AddrTuple { } } +#[derive(Default)] +struct StackState { + tuples: HashMap>, + closed: bool, +} + struct Shared { - tuples: RwLock>>, + state: RwLock, listening: RwLock>, tun: Arc, tuples_purge: broadcast::Sender, } +impl Shared { + fn is_closed(&self) -> bool { + self.state.read().unwrap().closed + } + + fn mark_closed_and_clear_tuples(&self) -> usize { + let mut state = self.state.write().unwrap(); + state.closed = true; + let len = state.tuples.len(); + state.tuples.clear(); + len + } +} + pub struct Stack { shared: Arc, local_ip: Ipv4Addr, @@ -353,7 +373,17 @@ impl Drop for Socket { fn drop(&mut self) { let tuple = AddrTuple::new(self.local_addr, self.remote_addr); // dissociates ourself from the dispatch map - assert!(self.shared.tuples.write().unwrap().remove(&tuple).is_some()); + let (removed, closed) = { + let mut state = self.shared.state.write().unwrap(); + (state.tuples.remove(&tuple).is_some(), state.closed) + }; + if !removed { + if closed { + trace!(?tuple, "Fake TCP tuple already removed after stack closed"); + } else { + warn!(?tuple, "Fake TCP tuple missing while dropping socket"); + } + } // purge cache let _ = self.shared.tuples_purge.send(tuple); @@ -400,7 +430,7 @@ impl Stack { ) -> Stack { let (tuples_purge_tx, _tuples_purge_rx) = broadcast::channel(16); let shared = Arc::new(Shared { - tuples: RwLock::new(HashMap::new()), + state: RwLock::new(StackState::default()), tun: tun.clone(), listening: RwLock::new(HashSet::new()), tuples_purge: tuples_purge_tx.clone(), @@ -426,19 +456,31 @@ impl Stack { self.shared.tun.driver_type() } + pub fn is_closed(&self) -> bool { + self.shared.is_closed() || self.reader_task.is_finished() + } + /// Listens for incoming connections on the given `port`. pub fn listen(&mut self, port: u16) { assert!(self.shared.listening.write().unwrap().insert(port)); } - pub async fn alloc_established_socket( - &mut self, + pub fn try_alloc_established_socket( + &self, local_addr: SocketAddr, remote_addr: SocketAddr, state: State, - ) -> Socket { + ) -> Option { let tuple = AddrTuple::new(local_addr, remote_addr); - let mut tuples = self.shared.tuples.write().unwrap(); + let mut stack_state = self.shared.state.write().unwrap(); + if stack_state.closed || self.reader_task.is_finished() { + stack_state.closed = true; + warn!( + ?tuple, + "fake_tcp stack is closed, refusing to allocate socket" + ); + return None; + } let (sock, incoming) = Socket::new( self.shared.clone(), // self.shared.tun.choose(&mut rng).unwrap().clone(), @@ -450,8 +492,8 @@ impl Stack { Some(0), // Initial ACK state, ); - assert!(tuples.insert(tuple, incoming).is_none()); - sock + assert!(stack_state.tuples.insert(tuple, incoming).is_none()); + Some(sock) } async fn reader_task( @@ -466,7 +508,22 @@ impl Stack { tokio::select! { size = tun.recv(&mut buf) => { - let size = size.unwrap(); + let size = match size { + Ok(size) => size, + Err(e) => { + let shared_tuple_count = shared.mark_closed_and_clear_tuples(); + let cached_tuple_count = tuples.len(); + tuples.clear(); + error!( + ?e, + driver_type = tun.driver_type(), + shared_tuple_count, + cached_tuple_count, + "fake_tcp tun recv failed, reader_task exiting" + ); + break; + } + }; tracing::trace!(len = size, ?buf, "PnetTun received packet"); let buf = buf.split().freeze(); @@ -494,8 +551,8 @@ impl Stack { } else { trace!("Cache miss, checking the shared tuples table for connection"); let sender = { - let tuples = shared.tuples.read().unwrap(); - tuples.get(&tuple).cloned() + let state = shared.state.read().unwrap(); + state.tuples.get(&tuple).cloned() }; if let Some(c) = sender { @@ -532,11 +589,107 @@ impl Stack { } }, tuple = tuples_purge.recv() => { - let tuple = tuple.unwrap(); - tuples.remove(&tuple); - trace!("Removed cached tuple: {:?}", tuple); + match tuple { + Ok(tuple) => { + tuples.remove(&tuple); + trace!("Removed cached tuple: {:?}", tuple); + } + Err(broadcast::error::RecvError::Lagged(skipped)) => { + let cached_tuple_count = tuples.len(); + tuples.clear(); + warn!( + skipped, + cached_tuple_count, + "fake_tcp tuples purge receiver lagged, cleared local cache" + ); + } + Err(broadcast::error::RecvError::Closed) => { + let shared_tuple_count = shared.mark_closed_and_clear_tuples(); + let cached_tuple_count = tuples.len(); + tuples.clear(); + warn!( + shared_tuple_count, + cached_tuple_count, + "fake_tcp tuples purge channel closed, reader_task exiting" + ); + break; + } + } } } } } } + +#[cfg(test)] +mod tests { + use super::*; + use std::io; + use tokio::{ + sync::Notify, + time::{Duration, timeout}, + }; + + #[derive(Default)] + struct FailingTun { + fail: Notify, + } + + impl FailingTun { + fn fail(&self) { + self.fail.notify_one(); + } + } + + #[async_trait::async_trait] + impl Tun for FailingTun { + async fn recv(&self, _packet: &mut BytesMut) -> Result { + self.fail.notified().await; + Err(io::Error::new(io::ErrorKind::BrokenPipe, "test tun closed")) + } + + fn try_send(&self, _packet: &Bytes) -> Result<(), io::Error> { + Ok(()) + } + + fn driver_type(&self) -> &'static str { + "test" + } + } + + #[tokio::test] + async fn reader_task_closes_sockets_on_tun_recv_error() { + let tun = Arc::new(FailingTun::default()); + let mut stack = Stack::new(tun.clone(), Ipv4Addr::LOCALHOST, None, None); + let socket = stack + .try_alloc_established_socket( + SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 10_000), + SocketAddr::new(Ipv4Addr::new(192, 0, 2, 1).into(), 20_000), + State::Established, + ) + .expect("socket allocation should succeed before tun failure"); + + tun.fail(); + + let join_result = timeout(Duration::from_secs(1), &mut stack.reader_task) + .await + .expect("reader task should exit after tun recv error"); + assert!(join_result.is_ok()); + assert!(stack.is_closed()); + + let mut buf = BytesMut::new(); + let recv_result = timeout(Duration::from_secs(1), socket.recv(&mut buf)) + .await + .expect("socket recv should not hang after reader task exits"); + assert_eq!(recv_result, None); + + let new_socket = stack.try_alloc_established_socket( + SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 10_001), + SocketAddr::new(Ipv4Addr::new(192, 0, 2, 1).into(), 20_001), + State::Established, + ); + assert!(new_socket.is_none()); + + drop(socket); + } +} diff --git a/easytier/src/tunnel/websocket.rs b/easytier/src/tunnel/websocket.rs index 1992fcfb..d7d1e6ae 100644 --- a/easytier/src/tunnel/websocket.rs +++ b/easytier/src/tunnel/websocket.rs @@ -118,6 +118,7 @@ impl WsTunnelListener { let (request, stream) = ServerBuilder::new() .limits(Limits::unlimited()) + .max_headers(128) .accept(stream) .await?; @@ -252,7 +253,8 @@ impl WsTunnelConnector { ), }; - let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap()); + let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap()) + .max_headers(128); let stream: MaybeTlsStream = if is_wss { init_crypto_provider(); let tls_conn =