Merge branch 'main' into tcp-stun

This commit is contained in:
fanyang
2026-06-23 23:47:32 +08:00
committed by GitHub
11 changed files with 656 additions and 177 deletions
Generated
+70 -13
View File
@@ -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",
@@ -119,7 +119,9 @@ fn sync_tun_event_receivers(receivers: &mut HashMap<String, EventBusSubscriber>)
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()
+3 -2
View File
@@ -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"
+150 -14
View File
@@ -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<ConfigSourceConfig>,
}
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<Mutex<Config>>,
@@ -597,11 +627,35 @@ impl TomlConfigLoader {
}
pub fn new_from_str(config_str: &str) -> Result<Self, anyhow::Error> {
let mut config = toml::de::from_str::<Config>(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<Self, anyhow::Error> {
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<Self, anyhow::Error> {
let mut config = toml::de::from_str::<Config>(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<Self, anyhow::Error> {
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<Self, anyhow::Error> {
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<String, serde_json::Value>,
) -> serde_json::Result<Flags> {
@@ -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("<unknown>"));
assert!(
error
.chain()
.any(|err| err.downcast_ref::<toml::de::Error>().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("<unknown>"));
}
#[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("<unknown>"));
}
#[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("<unknown>"));
}
#[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<String> = 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.
+8 -7
View File
@@ -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)?;
};
}
@@ -14,7 +14,7 @@ use std::{
};
use dashmap::{DashMap, DashSet};
use guarden::defer;
use guarden::{Guard, defer};
use tokio::{
sync::{
Mutex,
+3 -4
View File
@@ -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,
+130 -2
View File
@@ -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;
+58 -37
View File
@@ -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<tokio::net::TcpListener>,
// interface_name -> fake tcp stack
stack_map: DashMap<String, Arc<Mutex<stack::Stack>>>,
stack_map: DashMap<String, Arc<stack::Stack>>,
// 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<Arc<Mutex<stack::Stack>>, TunnelError> {
) -> Result<Arc<stack::Stack>, 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<Box<dyn Tunnel>, 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?;
+169 -16
View File
@@ -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<AddrTuple, flume::Sender<Bytes>>,
closed: bool,
}
struct Shared {
tuples: RwLock<HashMap<AddrTuple, flume::Sender<Bytes>>>,
state: RwLock<StackState>,
listening: RwLock<HashSet<u16>>,
tun: Arc<dyn Tun>,
tuples_purge: broadcast::Sender<AddrTuple>,
}
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<Shared>,
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<Socket> {
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<usize, io::Error> {
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);
}
}
+3 -1
View File
@@ -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<TcpStream> = if is_wss {
init_crypto_provider();
let tls_conn =