mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-08-06 12:39:51 +00:00
Merge branch 'main' into tcp-stun
This commit is contained in:
Generated
+70
-13
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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?;
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 =
|
||||
|
||||
Reference in New Issue
Block a user