mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-20 03:22:05 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
003fdefc63 | ||
|
|
a25c411249 | ||
|
|
0c11cefc04 | ||
|
|
8181830902 | ||
|
|
0a8c95879b |
Generated
+26
-72
@@ -241,16 +241,6 @@ 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"
|
||||
@@ -925,7 +915,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bfcfdc083699101d5a7965e49925975f2f55060f94f9a05e7187be95d530ca59"
|
||||
dependencies = [
|
||||
"once_cell",
|
||||
"proc-macro-crate 3.5.0",
|
||||
"proc-macro-crate 3.2.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
@@ -2244,7 +2234,6 @@ dependencies = [
|
||||
"aes-gcm",
|
||||
"anyhow",
|
||||
"arc-swap",
|
||||
"ariadne",
|
||||
"async-recursion",
|
||||
"async-ringbuf",
|
||||
"async-stream",
|
||||
@@ -2283,7 +2272,7 @@ dependencies = [
|
||||
"gethostname 0.5.0",
|
||||
"git-version",
|
||||
"globwalk",
|
||||
"guarden 0.2.0",
|
||||
"guarden",
|
||||
"hickory-client",
|
||||
"hickory-proto",
|
||||
"hickory-resolver",
|
||||
@@ -2329,7 +2318,7 @@ dependencies = [
|
||||
"prost-reflect-build",
|
||||
"prost-wkt-types",
|
||||
"quinn",
|
||||
"quinn-proto",
|
||||
"quinn-plaintext",
|
||||
"quote",
|
||||
"rand 0.8.5",
|
||||
"rcgen",
|
||||
@@ -2341,7 +2330,6 @@ dependencies = [
|
||||
"rstest",
|
||||
"rust-i18n",
|
||||
"rustls",
|
||||
"seahash",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serial_test",
|
||||
@@ -2469,7 +2457,7 @@ dependencies = [
|
||||
"dashmap",
|
||||
"easytier",
|
||||
"futures",
|
||||
"guarden 0.1.2",
|
||||
"guarden",
|
||||
"jsonwebtoken",
|
||||
"mimalloc",
|
||||
"mockall",
|
||||
@@ -3605,18 +3593,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ca87812d87fa82896df1adfb5c111cdeaae3edb6da028f5df002dcbd7df71454"
|
||||
dependencies = [
|
||||
"futures",
|
||||
"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",
|
||||
"guarden-macros",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
@@ -3631,18 +3608,6 @@ 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"
|
||||
@@ -5613,7 +5578,7 @@ version = "0.7.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "680998035259dcfcafe653688bf2aa6d3e2dc05e98be6ab46afb089dc84f1df8"
|
||||
dependencies = [
|
||||
"proc-macro-crate 2.0.0",
|
||||
"proc-macro-crate 3.2.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
@@ -6745,11 +6710,11 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proc-macro-crate"
|
||||
version = "3.5.0"
|
||||
version = "3.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f"
|
||||
checksum = "8ecf48c7ca261d60b74ab1a7b20da18bede46776b2e55535cb958eb595c5fa7b"
|
||||
dependencies = [
|
||||
"toml_edit 0.25.12+spec-1.1.0",
|
||||
"toml_edit 0.22.20",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -7057,6 +7022,18 @@ dependencies = [
|
||||
"web-time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quinn-plaintext"
|
||||
version = "0.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f3e617feaeb6493018fa35fc47ae8b630ac8903d8159e9e747018841b99bad3d"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"quinn-proto",
|
||||
"seahash",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quinn-proto"
|
||||
version = "0.11.12"
|
||||
@@ -7631,7 +7608,7 @@ checksum = "1f168d99749d307be9de54d23fd226628d99768225ef08f6ffb52e0182a27746"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"glob",
|
||||
"proc-macro-crate 3.5.0",
|
||||
"proc-macro-crate 3.2.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"regex",
|
||||
@@ -9955,7 +9932,8 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "tokio-websockets"
|
||||
version = "0.13.2"
|
||||
source = "git+https://github.com/EasyTier/tokio-websockets#dc9771c7c215882349c3cb328877550a3593df21"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "dad543404f98bfc969aeb71994105c592acfc6c43323fddcd016bb208d1c65cb"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
@@ -10030,15 +10008,6 @@ 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"
|
||||
@@ -10076,18 +10045,6 @@ 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"
|
||||
@@ -11935,9 +11892,6 @@ name = "winnow"
|
||||
version = "1.0.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "09dac053f1cd375980747450bfc7250c264eaae0583872e845c0c7cd578872b5"
|
||||
dependencies = [
|
||||
"memchr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "winreg"
|
||||
@@ -12319,7 +12273,7 @@ version = "5.14.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "897e79616e84aac4b2c46e9132a4f63b93105d54fe8c0e8f6bffc21fa8d49222"
|
||||
dependencies = [
|
||||
"proc-macro-crate 3.5.0",
|
||||
"proc-macro-crate 3.2.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
@@ -12556,7 +12510,7 @@ version = "5.10.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5b59b012ebe9c46656f9cc08d8da8b4c726510aef12559da3e5f1bf72780752c"
|
||||
dependencies = [
|
||||
"proc-macro-crate 3.5.0",
|
||||
"proc-macro-crate 3.2.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
|
||||
@@ -118,7 +118,11 @@ pub(crate) fn set_tun_fd(
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn collect_runtime_state() -> RuntimeAggregateState {
|
||||
pub(crate) fn get_runtime_snapshot() -> RuntimeAggregateState {
|
||||
get_runtime_snapshot_inner()
|
||||
}
|
||||
|
||||
pub(crate) fn get_runtime_snapshot_inner() -> RuntimeAggregateState {
|
||||
let infos = match ASYNC_RUNTIME.block_on(INSTANCE_MANAGER.collect_network_infos()) {
|
||||
Ok(infos) => infos,
|
||||
Err(err) => {
|
||||
|
||||
@@ -3,4 +3,6 @@ mod routing;
|
||||
mod socket_server;
|
||||
|
||||
pub(crate) use routing::aggregate_requested_tun_routes;
|
||||
pub use socket_server::{start_local_socket_server, stop_local_socket_server};
|
||||
pub use socket_server::{
|
||||
set_snapshot_broadcast_enabled, start_local_socket_server, stop_local_socket_server,
|
||||
};
|
||||
|
||||
@@ -32,13 +32,6 @@ pub(crate) fn send_local_socket_message(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn shrink_clients_if_sparse(clients: &mut Vec<UnixStream>) {
|
||||
let sparse_limit = clients.len().saturating_mul(2).max(4);
|
||||
if clients.capacity() > sparse_limit {
|
||||
clients.shrink_to_fit();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn broadcast_local_socket_message(
|
||||
clients: &mut Vec<UnixStream>,
|
||||
message_type: &str,
|
||||
@@ -52,7 +45,6 @@ pub(crate) fn broadcast_local_socket_message(
|
||||
active_clients.push(client);
|
||||
}
|
||||
}
|
||||
shrink_clients_if_sparse(&mut active_clients);
|
||||
*clients = active_clients;
|
||||
delivered
|
||||
}
|
||||
@@ -87,7 +79,6 @@ pub(crate) fn broadcast_local_socket_json_payload_message(
|
||||
active_clients.push(client);
|
||||
}
|
||||
}
|
||||
shrink_clients_if_sparse(&mut active_clients);
|
||||
*clients = active_clients;
|
||||
delivered
|
||||
}
|
||||
|
||||
@@ -1,20 +1,13 @@
|
||||
use super::protocol::{
|
||||
TunRequestPayload, broadcast_local_socket_json_payload_message, broadcast_local_socket_message,
|
||||
};
|
||||
use crate::collect_runtime_state_inner;
|
||||
use crate::INSTANCE_MANAGER;
|
||||
use crate::config::repository::kernel_socket_path;
|
||||
use crate::get_runtime_snapshot_inner;
|
||||
use crate::kernel_bridge::routing::aggregate_tun_routes;
|
||||
use crate::runtime::state::runtime_state::{
|
||||
PeerConnInfo as RuntimePeerConnInfo, RuntimeAggregateState, peer_conn_to_view,
|
||||
};
|
||||
use crate::{ASYNC_RUNTIME, INSTANCE_MANAGER};
|
||||
use easytier::common::global_ctx::{EventBusSubscriber, GlobalCtxEvent};
|
||||
use easytier::proto::api::instance::ListPeerRequest;
|
||||
use easytier::proto::rpc_types::controller::BaseController;
|
||||
use once_cell::sync::Lazy;
|
||||
use serde::Serialize;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::hash::Hash;
|
||||
use std::io::ErrorKind;
|
||||
use std::os::unix::net::{UnixListener, UnixStream};
|
||||
use std::path::PathBuf;
|
||||
@@ -30,75 +23,13 @@ struct LocalSocketState {
|
||||
}
|
||||
|
||||
static LOCAL_SOCKET_STATE: Lazy<Mutex<Option<LocalSocketState>>> = Lazy::new(|| Mutex::new(None));
|
||||
static SNAPSHOT_BROADCAST_ENABLED: AtomicBool = AtomicBool::new(true);
|
||||
const SOCKET_TICK_INTERVAL: Duration = Duration::from_millis(250);
|
||||
const TRAFFIC_STATS_INTERVAL: Duration = Duration::from_secs(1);
|
||||
const INSTANCE_POLL_INTERVAL: Duration = Duration::from_secs(1);
|
||||
const TUN_FAST_CHECK_WINDOW: Duration = Duration::from_secs(8);
|
||||
const EVENT_RECEIVER_SYNC_INTERVAL: Duration = Duration::from_secs(1);
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct TrafficStatsPayload {
|
||||
instances: Vec<InstanceTrafficStats>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct InstanceTrafficStats {
|
||||
config_id: String,
|
||||
instance_id: String,
|
||||
rx_bytes: i64,
|
||||
tx_bytes: i64,
|
||||
peers: Vec<PeerTrafficStats>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct PeerTrafficStats {
|
||||
peer_id: i64,
|
||||
rx_bytes: i64,
|
||||
tx_bytes: i64,
|
||||
total_bytes: i64,
|
||||
latency_us: i64,
|
||||
loss_rate: f64,
|
||||
}
|
||||
|
||||
struct PendingPeerEvent {
|
||||
event: &'static str,
|
||||
instance_id: String,
|
||||
peer_id: i64,
|
||||
conn: Option<RuntimePeerConnInfo>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct DrainedKernelEvents {
|
||||
tun_refresh: bool,
|
||||
topology_lost: bool,
|
||||
peer_events: Vec<PendingPeerEvent>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
struct RuntimePeerEventPayload {
|
||||
event: &'static str,
|
||||
config_id: String,
|
||||
instance_id: String,
|
||||
peer_id: i64,
|
||||
conn: Option<RuntimePeerConnInfo>,
|
||||
}
|
||||
|
||||
fn shrink_hash_map_if_sparse<K: Eq + Hash, V>(map: &mut HashMap<K, V>) {
|
||||
let sparse_limit = map.len().saturating_mul(2).max(8);
|
||||
if map.capacity() > sparse_limit {
|
||||
map.shrink_to_fit();
|
||||
}
|
||||
}
|
||||
|
||||
fn shrink_hash_set_if_sparse<T: Eq + Hash>(set: &mut HashSet<T>) {
|
||||
let sparse_limit = set.len().saturating_mul(2).max(8);
|
||||
if set.capacity() > sparse_limit {
|
||||
set.shrink_to_fit();
|
||||
}
|
||||
pub fn set_snapshot_broadcast_enabled(enabled: bool) {
|
||||
SNAPSHOT_BROADCAST_ENABLED.store(enabled, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
fn sync_tun_event_receivers(receivers: &mut HashMap<String, EventBusSubscriber>) {
|
||||
@@ -113,67 +44,36 @@ fn sync_tun_event_receivers(receivers: &mut HashMap<String, EventBusSubscriber>)
|
||||
}
|
||||
}
|
||||
receivers.retain(|instance_id, _| active_instance_ids.contains(instance_id));
|
||||
shrink_hash_map_if_sparse(receivers);
|
||||
}
|
||||
|
||||
fn event_needs_tun_refresh(event: &GlobalCtxEvent) -> bool {
|
||||
matches!(
|
||||
event,
|
||||
GlobalCtxEvent::DhcpIpv4Changed(_, _)
|
||||
| GlobalCtxEvent::ProxyCidrsUpdated(_, _)
|
||||
| GlobalCtxEvent::DhcpIpv4Conflicted(_)
|
||||
| GlobalCtxEvent::PublicIpv6Changed(_, _)
|
||||
| GlobalCtxEvent::PublicIpv6RoutesUpdated(_, _)
|
||||
| GlobalCtxEvent::ProxyCidrsUpdated(_, _)
|
||||
| GlobalCtxEvent::ConfigPatched(_)
|
||||
| GlobalCtxEvent::PeerAdded(_)
|
||||
| GlobalCtxEvent::PeerRemoved(_)
|
||||
| GlobalCtxEvent::PeerConnAdded(_)
|
||||
| GlobalCtxEvent::PeerConnRemoved(_)
|
||||
)
|
||||
}
|
||||
|
||||
fn drain_kernel_events(receivers: &mut HashMap<String, EventBusSubscriber>) -> DrainedKernelEvents {
|
||||
let mut drained = DrainedKernelEvents::default();
|
||||
fn drain_tun_refresh_events(receivers: &mut HashMap<String, EventBusSubscriber>) -> bool {
|
||||
let mut refresh_needed = false;
|
||||
let mut closed_receivers = Vec::new();
|
||||
for (instance_id, receiver) in receivers.iter_mut() {
|
||||
loop {
|
||||
match receiver.try_recv() {
|
||||
Ok(event) => {
|
||||
drained.tun_refresh = event_needs_tun_refresh(&event) || drained.tun_refresh;
|
||||
match event {
|
||||
GlobalCtxEvent::PeerAdded(peer_id) => {
|
||||
drained.peer_events.push(PendingPeerEvent {
|
||||
event: "peer_added",
|
||||
instance_id: instance_id.clone(),
|
||||
peer_id: peer_id as i64,
|
||||
conn: None,
|
||||
});
|
||||
}
|
||||
GlobalCtxEvent::PeerRemoved(peer_id) => {
|
||||
drained.peer_events.push(PendingPeerEvent {
|
||||
event: "peer_removed",
|
||||
instance_id: instance_id.clone(),
|
||||
peer_id: peer_id as i64,
|
||||
conn: None,
|
||||
});
|
||||
}
|
||||
GlobalCtxEvent::PeerConnAdded(conn_info) => {
|
||||
let peer_id = conn_info.peer_id as i64;
|
||||
drained.peer_events.push(PendingPeerEvent {
|
||||
event: "peer_conn_added",
|
||||
instance_id: instance_id.clone(),
|
||||
peer_id,
|
||||
conn: Some(peer_conn_to_view(conn_info)),
|
||||
});
|
||||
}
|
||||
GlobalCtxEvent::PeerConnRemoved(conn_info) => {
|
||||
let peer_id = conn_info.peer_id as i64;
|
||||
drained.peer_events.push(PendingPeerEvent {
|
||||
event: "peer_conn_removed",
|
||||
instance_id: instance_id.clone(),
|
||||
peer_id,
|
||||
conn: Some(peer_conn_to_view(conn_info)),
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
refresh_needed = event_needs_tun_refresh(&event) || refresh_needed;
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::TryRecvError::Empty) => break,
|
||||
Err(tokio::sync::broadcast::error::TryRecvError::Lagged(_)) => {
|
||||
drained.topology_lost = true;
|
||||
refresh_needed = true;
|
||||
continue;
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::TryRecvError::Closed) => {
|
||||
@@ -186,124 +86,7 @@ fn drain_kernel_events(receivers: &mut HashMap<String, EventBusSubscriber>) -> D
|
||||
for instance_id in closed_receivers {
|
||||
receivers.remove(&instance_id);
|
||||
}
|
||||
drained
|
||||
}
|
||||
|
||||
fn broadcast_runtime_peer_events(
|
||||
clients: &mut Vec<UnixStream>,
|
||||
peer_events: Vec<PendingPeerEvent>,
|
||||
) {
|
||||
for event in peer_events {
|
||||
let payload = RuntimePeerEventPayload {
|
||||
event: event.event,
|
||||
config_id: event.instance_id.clone(),
|
||||
instance_id: event.instance_id,
|
||||
peer_id: event.peer_id,
|
||||
conn: event.conn,
|
||||
};
|
||||
match serde_json::to_string(&payload) {
|
||||
Ok(json) => {
|
||||
let _ = broadcast_local_socket_json_payload_message(
|
||||
clients,
|
||||
"runtime_peer_event",
|
||||
&json,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
ohrs_log_error!("[Rust] serialize runtime peer event failed: {}", err);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn tun_candidate_ids(snapshot: &RuntimeAggregateState) -> HashSet<String> {
|
||||
snapshot
|
||||
.instances
|
||||
.iter()
|
||||
.filter(|instance| instance.running && instance.tun_required)
|
||||
.map(|instance| instance.instance_id.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn collect_traffic_stats() -> TrafficStatsPayload {
|
||||
let services = INSTANCE_MANAGER
|
||||
.iter()
|
||||
.filter_map(|instance| {
|
||||
instance
|
||||
.value()
|
||||
.get_api_service()
|
||||
.map(|api_service| (instance.key().to_string(), api_service))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let instances = ASYNC_RUNTIME.block_on(async {
|
||||
let mut instances = Vec::new();
|
||||
for (instance_id, api_service) in services {
|
||||
let peers = match api_service
|
||||
.get_peer_manage_service()
|
||||
.list_peer(BaseController::default(), ListPeerRequest::default())
|
||||
.await
|
||||
{
|
||||
Ok(response) => response.peer_infos,
|
||||
Err(err) => {
|
||||
ohrs_log_debug!(
|
||||
"[Rust] collect traffic stats list_peer failed instance={}: {}",
|
||||
instance_id,
|
||||
err
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let mut instance_rx_bytes = 0i64;
|
||||
let mut instance_tx_bytes = 0i64;
|
||||
let mut peer_stats = Vec::with_capacity(peers.len());
|
||||
|
||||
for peer in peers {
|
||||
let mut peer_rx_bytes = 0i64;
|
||||
let mut peer_tx_bytes = 0i64;
|
||||
let mut latency_us = i64::MAX;
|
||||
let mut loss_rate = 0f64;
|
||||
|
||||
for conn in peer.conns {
|
||||
if let Some(stats) = conn.stats {
|
||||
let rx_bytes = stats.rx_bytes as i64;
|
||||
let tx_bytes = stats.tx_bytes as i64;
|
||||
peer_rx_bytes += rx_bytes;
|
||||
peer_tx_bytes += tx_bytes;
|
||||
latency_us = latency_us.min(stats.latency_us as i64);
|
||||
}
|
||||
loss_rate = loss_rate.max(conn.loss_rate as f64);
|
||||
}
|
||||
|
||||
instance_rx_bytes += peer_rx_bytes;
|
||||
instance_tx_bytes += peer_tx_bytes;
|
||||
peer_stats.push(PeerTrafficStats {
|
||||
peer_id: peer.peer_id as i64,
|
||||
rx_bytes: peer_rx_bytes,
|
||||
tx_bytes: peer_tx_bytes,
|
||||
total_bytes: peer_rx_bytes + peer_tx_bytes,
|
||||
latency_us: if latency_us == i64::MAX {
|
||||
-1
|
||||
} else {
|
||||
latency_us
|
||||
},
|
||||
loss_rate,
|
||||
});
|
||||
}
|
||||
|
||||
instances.push(InstanceTrafficStats {
|
||||
config_id: instance_id.clone(),
|
||||
instance_id,
|
||||
rx_bytes: instance_rx_bytes,
|
||||
tx_bytes: instance_tx_bytes,
|
||||
peers: peer_stats,
|
||||
});
|
||||
}
|
||||
instances
|
||||
});
|
||||
|
||||
TrafficStatsPayload { instances }
|
||||
refresh_needed
|
||||
}
|
||||
|
||||
pub fn start_local_socket_server() -> bool {
|
||||
@@ -348,25 +131,21 @@ pub fn start_local_socket_server() -> bool {
|
||||
let stop_flag = std::sync::Arc::new(AtomicBool::new(false));
|
||||
let worker_stop_flag = stop_flag.clone();
|
||||
let worker = thread::spawn(move || {
|
||||
let mut last_topology_json = String::new();
|
||||
let mut last_snapshot_json = String::new();
|
||||
let mut delivered_tun_requests = HashSet::new();
|
||||
let mut last_tun_route_signatures = HashMap::<String, String>::new();
|
||||
let mut tun_fast_until = Instant::now() + TUN_FAST_CHECK_WINDOW;
|
||||
let mut tun_bootstrap_done = false;
|
||||
let mut last_event_receiver_sync_at: Option<Instant> = None;
|
||||
let mut last_traffic_stats_at: Option<Instant> = None;
|
||||
let mut last_instance_poll_at: Option<Instant> = None;
|
||||
let mut tun_event_receivers = HashMap::<String, EventBusSubscriber>::new();
|
||||
let mut clients = Vec::<UnixStream>::new();
|
||||
|
||||
while !worker_stop_flag.load(Ordering::Relaxed) {
|
||||
let mut full_topology_dirty = false;
|
||||
let mut accepted_client = false;
|
||||
loop {
|
||||
match listener.accept() {
|
||||
Ok((stream, _addr)) => {
|
||||
accepted_client = true;
|
||||
full_topology_dirty = true;
|
||||
clients.push(stream);
|
||||
tun_fast_until = Instant::now() + TUN_FAST_CHECK_WINDOW;
|
||||
tun_bootstrap_done = false;
|
||||
@@ -379,21 +158,15 @@ pub fn start_local_socket_server() -> bool {
|
||||
}
|
||||
}
|
||||
|
||||
let snapshot_enabled = SNAPSHOT_BROADCAST_ENABLED.load(Ordering::Relaxed);
|
||||
if clients.is_empty() {
|
||||
if !last_topology_json.is_empty() {
|
||||
last_topology_json.clear();
|
||||
last_topology_json.shrink_to_fit();
|
||||
if !last_snapshot_json.is_empty() {
|
||||
last_snapshot_json.clear();
|
||||
}
|
||||
delivered_tun_requests.clear();
|
||||
shrink_hash_set_if_sparse(&mut delivered_tun_requests);
|
||||
last_tun_route_signatures.clear();
|
||||
shrink_hash_map_if_sparse(&mut last_tun_route_signatures);
|
||||
tun_event_receivers.clear();
|
||||
shrink_hash_map_if_sparse(&mut tun_event_receivers);
|
||||
clients.shrink_to_fit();
|
||||
last_event_receiver_sync_at = None;
|
||||
last_traffic_stats_at = None;
|
||||
last_instance_poll_at = None;
|
||||
tun_bootstrap_done = false;
|
||||
thread::sleep(SOCKET_TICK_INTERVAL);
|
||||
continue;
|
||||
@@ -408,143 +181,115 @@ pub fn start_local_socket_server() -> bool {
|
||||
sync_tun_event_receivers(&mut tun_event_receivers);
|
||||
last_event_receiver_sync_at = Some(now);
|
||||
}
|
||||
let drained_events = drain_kernel_events(&mut tun_event_receivers);
|
||||
let tun_refresh = drained_events.tun_refresh;
|
||||
let topology_lost = drained_events.topology_lost;
|
||||
let peer_events = drained_events.peer_events;
|
||||
if topology_lost {
|
||||
full_topology_dirty = true;
|
||||
}
|
||||
if tun_refresh {
|
||||
if drain_tun_refresh_events(&mut tun_event_receivers) {
|
||||
tun_bootstrap_done = false;
|
||||
tun_fast_until = now + TUN_FAST_CHECK_WINDOW;
|
||||
}
|
||||
if !peer_events.is_empty() {
|
||||
broadcast_runtime_peer_events(&mut clients, peer_events);
|
||||
}
|
||||
let should_collect_traffic_stats = last_traffic_stats_at
|
||||
.map(|last| now.duration_since(last) >= TRAFFIC_STATS_INTERVAL)
|
||||
.unwrap_or(true);
|
||||
if should_collect_traffic_stats {
|
||||
last_traffic_stats_at = Some(now);
|
||||
match serde_json::to_string(&collect_traffic_stats()) {
|
||||
Ok(json) => {
|
||||
let _ = broadcast_local_socket_json_payload_message(
|
||||
&mut clients,
|
||||
"traffic_stats",
|
||||
&json,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
ohrs_log_error!("[Rust] serialize traffic stats failed: {}", err);
|
||||
}
|
||||
}
|
||||
}
|
||||
let should_poll_instance = last_instance_poll_at
|
||||
.map(|last| now.duration_since(last) >= INSTANCE_POLL_INTERVAL)
|
||||
.unwrap_or(true);
|
||||
let should_collect_topology = accepted_client
|
||||
|| full_topology_dirty
|
||||
|| tun_refresh
|
||||
|| should_poll_instance
|
||||
let should_collect_snapshot = snapshot_enabled
|
||||
|| accepted_client
|
||||
|| (!tun_bootstrap_done && now < tun_fast_until);
|
||||
if !should_collect_topology {
|
||||
if !should_collect_snapshot {
|
||||
if !last_snapshot_json.is_empty() {
|
||||
last_snapshot_json.clear();
|
||||
}
|
||||
thread::sleep(SOCKET_TICK_INTERVAL);
|
||||
continue;
|
||||
}
|
||||
|
||||
let snapshot = collect_runtime_state_inner();
|
||||
last_instance_poll_at = Some(now);
|
||||
match serde_json::to_string(&snapshot) {
|
||||
Ok(json) => {
|
||||
if accepted_client || full_topology_dirty || json != last_topology_json {
|
||||
let _ = broadcast_local_socket_json_payload_message(
|
||||
&mut clients,
|
||||
"runtime_topology",
|
||||
&json,
|
||||
);
|
||||
last_topology_json = json;
|
||||
let snapshot = get_runtime_snapshot_inner();
|
||||
if snapshot_enabled {
|
||||
let snapshot_json = match serde_json::to_string(&snapshot) {
|
||||
Ok(json) => json,
|
||||
Err(err) => {
|
||||
ohrs_log_error!("[Rust] serialize runtime snapshot failed: {}", err);
|
||||
thread::sleep(SOCKET_TICK_INTERVAL);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if accepted_client || snapshot_json != last_snapshot_json {
|
||||
let _ = broadcast_local_socket_json_payload_message(
|
||||
&mut clients,
|
||||
"runtime_snapshot",
|
||||
&snapshot_json,
|
||||
);
|
||||
last_snapshot_json = snapshot_json;
|
||||
}
|
||||
Err(err) => {
|
||||
ohrs_log_error!("[Rust] serialize runtime topology failed: {}", err);
|
||||
}
|
||||
} else if !last_snapshot_json.is_empty() {
|
||||
last_snapshot_json.clear();
|
||||
}
|
||||
|
||||
let active_tun_candidate_ids = tun_candidate_ids(&snapshot);
|
||||
delivered_tun_requests
|
||||
.retain(|instance_id| active_tun_candidate_ids.contains(instance_id));
|
||||
last_tun_route_signatures
|
||||
.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 mut saw_running_instance = false;
|
||||
let mut saw_tun_candidate = false;
|
||||
for instance in snapshot.instances.iter() {
|
||||
if instance.running {
|
||||
saw_running_instance = 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);
|
||||
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 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 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 = accepted_client
|
||||
|| !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);
|
||||
}
|
||||
} else {
|
||||
delivered_tun_requests.remove(&instance.instance_id);
|
||||
last_tun_route_signatures.remove(&instance.instance_id);
|
||||
}
|
||||
}
|
||||
if !delivered_tun_requests.is_empty()
|
||||
|| (saw_running_instance && !saw_tun_candidate)
|
||||
|| now >= tun_fast_until
|
||||
if !snapshot_enabled
|
||||
&& (!delivered_tun_requests.is_empty()
|
||||
|| (saw_running_instance && !saw_tun_candidate)
|
||||
|| now >= tun_fast_until)
|
||||
{
|
||||
tun_bootstrap_done = true;
|
||||
}
|
||||
|
||||
@@ -63,7 +63,7 @@ use easytier::proto::api::manage::NetworkConfig;
|
||||
use easytier::proto::api::manage::NetworkingMethod;
|
||||
use easytier::web_client::{WebClient, WebClientHooks, run_web_client};
|
||||
use kernel_bridge::{
|
||||
start_local_socket_server as start_local_socket_server_inner,
|
||||
set_snapshot_broadcast_enabled, start_local_socket_server as start_local_socket_server_inner,
|
||||
stop_local_socket_server as stop_local_socket_server_inner,
|
||||
};
|
||||
use napi_derive_ohos::napi;
|
||||
@@ -517,8 +517,18 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn collect_runtime_state_inner() -> RuntimeAggregateState {
|
||||
exports::runtime_api::collect_runtime_state()
|
||||
#[napi]
|
||||
pub fn get_runtime_snapshot() -> RuntimeAggregateState {
|
||||
exports::runtime_api::get_runtime_snapshot()
|
||||
}
|
||||
|
||||
#[napi]
|
||||
pub fn set_kernel_snapshot_enabled(enabled: bool) {
|
||||
set_snapshot_broadcast_enabled(enabled);
|
||||
}
|
||||
|
||||
pub(crate) fn get_runtime_snapshot_inner() -> RuntimeAggregateState {
|
||||
exports::runtime_api::get_runtime_snapshot_inner()
|
||||
}
|
||||
|
||||
#[napi]
|
||||
|
||||
@@ -324,7 +324,7 @@ fn route_to_view(route: api::instance::Route) -> RouteView {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn peer_conn_to_view(conn: api::instance::PeerConnInfo) -> PeerConnInfo {
|
||||
fn peer_conn_to_view(conn: api::instance::PeerConnInfo) -> PeerConnInfo {
|
||||
let stats = conn.stats.map(|stats| PeerConnStats {
|
||||
rx_bytes: stats.rx_bytes as i64,
|
||||
tx_bytes: stats.tx_bytes as i64,
|
||||
|
||||
+4
-6
@@ -51,7 +51,7 @@ time = "0.3"
|
||||
toml = "0.8.12"
|
||||
chrono = { version = "0.4.37", features = ["serde"] }
|
||||
|
||||
guarden = "0.2"
|
||||
guarden = "0.1"
|
||||
|
||||
delegate = "0.13.5"
|
||||
|
||||
@@ -82,8 +82,7 @@ pin-project-lite = "0.2.13"
|
||||
atomic_refcell = "0.1.13"
|
||||
|
||||
quinn = { version = "0.11.8", optional = true, features = ["ring"] }
|
||||
quinn-proto = { version = "0.11.12", optional = true }
|
||||
seahash = { version = "4.1.0", optional = true }
|
||||
quinn-plaintext = { version = "0.3.0", optional = true }
|
||||
|
||||
rustls = { version = "0.23.0", features = [
|
||||
"ring", "tls12"
|
||||
@@ -91,7 +90,7 @@ rustls = { version = "0.23.0", features = [
|
||||
rcgen = { version = "0.12.1", optional = true }
|
||||
|
||||
# for websocket
|
||||
tokio-websockets = { version = "0.13.2", git = "https://github.com/EasyTier/tokio-websockets", optional = true, features = [
|
||||
tokio-websockets = { version = "0.13.2", optional = true, features = [
|
||||
"rustls-webpki-roots",
|
||||
"client",
|
||||
"server",
|
||||
@@ -134,7 +133,6 @@ 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"
|
||||
@@ -375,7 +373,7 @@ full = [
|
||||
"zstd",
|
||||
]
|
||||
wireguard = ["dep:boringtun", "dep:ring"]
|
||||
quic = ["dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "dep:rcgen"]
|
||||
quic = ["dep:quinn", "dep:quinn-plaintext", "dep:rustls", "dep:rcgen"]
|
||||
kcp = ["dep:kcp-sys"]
|
||||
mimalloc = ["dep:mimalloc"]
|
||||
aes-gcm = ["dep:aes-gcm"]
|
||||
|
||||
@@ -49,8 +49,8 @@ core_clap:
|
||||
en: "manually specify the public IPv6 subnet to share, instead of auto-detecting from system routes"
|
||||
zh-CN: "手动指定要共享的公网 IPv6 子网,不自动从系统路由检测"
|
||||
dhcp:
|
||||
en: "automatically determine and set IP address by Easytier, and the IP address starts from 10.0.0.1 by default. Warning, if there is an IP conflict in the network when using DHCP, the IP will be automatically changed."
|
||||
zh-CN: "由Easytier自动确定并设置IP地址,默认从10.0.0.1开始。警告:在使用DHCP时,如果网络中出现IP冲突,IP将自动更改。"
|
||||
en: "automatically determine and set IP address by Easytier. The subnet is derived from a connected peer's IPv4 or defaults to 10.126.126.0/24. Warning, if there is an IP conflict in the network when using DHCP, the IP will be automatically changed. Optionally specify a CIDR subnet (e.g. -d 10.0.0.0/24, prefix <= /30) to pin the DHCP address range."
|
||||
zh-CN: "由Easytier自动确定并设置IP地址。子网从已连接对等节点的IPv4派生,或默认使用10.126.126.0/24。警告:在使用DHCP时,如果网络中出现IP冲突,IP将自动更改。可选指定CIDR子网(如 -d 10.0.0.0/24,前缀 <= /30)来固定DHCP地址范围。"
|
||||
peers:
|
||||
en: "peers to connect initially"
|
||||
zh-CN: "最初要连接的对等节点"
|
||||
|
||||
+35
-159
@@ -6,7 +6,6 @@ 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;
|
||||
@@ -186,6 +185,9 @@ pub trait ConfigLoader: Send + Sync {
|
||||
fn get_dhcp(&self) -> bool;
|
||||
fn set_dhcp(&self, dhcp: bool);
|
||||
|
||||
fn get_dhcp_cidr(&self) -> Option<cidr::Ipv4Cidr>;
|
||||
fn set_dhcp_cidr(&self, cidr: Option<cidr::Ipv4Cidr>);
|
||||
|
||||
fn add_proxy_cidr(
|
||||
&self,
|
||||
cidr: cidr::Ipv4Cidr,
|
||||
@@ -536,6 +538,7 @@ struct Config {
|
||||
ipv6_public_addr_auto: Option<bool>,
|
||||
ipv6_public_addr_prefix: Option<String>,
|
||||
dhcp: Option<bool>,
|
||||
dhcp_cidr: Option<String>,
|
||||
network_identity: Option<NetworkIdentity>,
|
||||
listeners: Option<Vec<url::Url>>,
|
||||
mapped_listeners: Option<Vec<url::Url>>,
|
||||
@@ -570,35 +573,6 @@ 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>>,
|
||||
@@ -621,39 +595,12 @@ impl TomlConfigLoader {
|
||||
}
|
||||
|
||||
pub fn new_from_str(config_str: &str) -> Result<Self, anyhow::Error> {
|
||||
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)
|
||||
})?;
|
||||
let mut config = toml::de::from_str::<Config>(config_str)
|
||||
.with_context(|| format!("failed to parse config file: {}", config_str))?;
|
||||
|
||||
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")?,
|
||||
);
|
||||
config.flags_struct = Some(Self::gen_flags(config.flags.clone().unwrap_or_default()));
|
||||
let has_network_identity = config.network_identity.is_some();
|
||||
|
||||
let config = TomlConfigLoader {
|
||||
@@ -685,15 +632,21 @@ impl TomlConfigLoader {
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
fn gen_flags(
|
||||
flags_hashmap: serde_json::Map<String, serde_json::Value>,
|
||||
) -> serde_json::Result<Flags> {
|
||||
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>) -> Flags {
|
||||
let mut merged_hashmap = match serde_json::to_value(gen_default_flags()) {
|
||||
Ok(serde_json::Value::Object(map)) => map,
|
||||
_ => serde_json::Map::new(),
|
||||
};
|
||||
merged_hashmap.extend(flags_hashmap);
|
||||
serde_json::from_value(serde_json::Value::Object(merged_hashmap))
|
||||
serde_json::from_value(serde_json::Value::Object(merged_hashmap)).unwrap()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -812,13 +765,26 @@ impl ConfigLoader for TomlConfigLoader {
|
||||
}
|
||||
|
||||
fn get_dhcp(&self) -> bool {
|
||||
self.config.lock().unwrap().dhcp.unwrap_or_default()
|
||||
let config = self.config.lock().unwrap();
|
||||
config.dhcp.unwrap_or_default() || config.dhcp_cidr.is_some()
|
||||
}
|
||||
|
||||
fn set_dhcp(&self, dhcp: bool) {
|
||||
self.config.lock().unwrap().dhcp = Some(dhcp);
|
||||
}
|
||||
|
||||
fn get_dhcp_cidr(&self) -> Option<cidr::Ipv4Cidr> {
|
||||
let locked_config = self.config.lock().unwrap();
|
||||
locked_config
|
||||
.dhcp_cidr
|
||||
.as_ref()
|
||||
.and_then(|s| s.parse().ok())
|
||||
}
|
||||
|
||||
fn set_dhcp_cidr(&self, cidr: Option<cidr::Ipv4Cidr>) {
|
||||
self.config.lock().unwrap().dhcp_cidr = cidr.map(|c| c.to_string());
|
||||
}
|
||||
|
||||
fn add_proxy_cidr(
|
||||
&self,
|
||||
cidr: cidr::Ipv4Cidr,
|
||||
@@ -1250,13 +1216,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_with_source("stdin", &stdin)?;
|
||||
let config = TomlConfigLoader::new_from_str(&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.display()))?;
|
||||
.with_context(|| format!("failed to read config file: {:?}", config_file))?;
|
||||
|
||||
let (expanded_config_str, uses_env_vars) = if disable_env_parsing {
|
||||
(config_str.clone(), false)
|
||||
@@ -1278,8 +1244,8 @@ pub async fn load_config_from_file(
|
||||
);
|
||||
}
|
||||
|
||||
let source_name = config_file.display().to_string();
|
||||
let config = TomlConfigLoader::new_from_str_with_source(&source_name, &expanded_config_str)?;
|
||||
let config = TomlConfigLoader::new_from_str(&expanded_config_str)
|
||||
.with_context(|| format!("failed to load config file: {:?}", config_file))?;
|
||||
|
||||
let mut control = ConfigFileControl::from_path(config_file.clone()).await;
|
||||
|
||||
@@ -1319,96 +1285,6 @@ 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.
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
use dashmap::DashMap;
|
||||
use parking_lot::Mutex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::cell::UnsafeCell;
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::time::interval;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
@@ -23,8 +22,6 @@ pub enum MetricName {
|
||||
PeerRpcDuration,
|
||||
/// RPC errors
|
||||
PeerRpcErrors,
|
||||
/// RPC/control packets dropped because the peer RPC queue is unavailable
|
||||
PeerRpcPacketQueueDrops,
|
||||
|
||||
/// Data-plane traffic bytes sent
|
||||
TrafficBytesTx,
|
||||
@@ -118,7 +115,6 @@ impl fmt::Display for MetricName {
|
||||
MetricName::PeerRpcServerRx => write!(f, "peer_rpc_server_rx"),
|
||||
MetricName::PeerRpcDuration => write!(f, "peer_rpc_duration_ms"),
|
||||
MetricName::PeerRpcErrors => write!(f, "peer_rpc_errors"),
|
||||
MetricName::PeerRpcPacketQueueDrops => write!(f, "peer_rpc_packet_queue_drops"),
|
||||
|
||||
MetricName::TrafficBytesTx => write!(f, "traffic_bytes_tx"),
|
||||
MetricName::TrafficBytesTxByInstance => write!(f, "traffic_bytes_tx_by_instance"),
|
||||
@@ -378,10 +374,10 @@ impl Default for LabelSet {
|
||||
}
|
||||
}
|
||||
|
||||
/// UnsafeCounter provides a high-performance atomic counter
|
||||
/// UnsafeCounter provides a high-performance counter using UnsafeCell
|
||||
#[derive(Debug)]
|
||||
pub struct UnsafeCounter {
|
||||
value: AtomicU64,
|
||||
value: UnsafeCell<u64>,
|
||||
}
|
||||
|
||||
impl Default for UnsafeCounter {
|
||||
@@ -393,79 +389,121 @@ impl Default for UnsafeCounter {
|
||||
impl UnsafeCounter {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
value: AtomicU64::new(0),
|
||||
value: UnsafeCell::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_with_value(initial: u64) -> Self {
|
||||
Self {
|
||||
value: AtomicU64::new(initial),
|
||||
value: UnsafeCell::new(initial),
|
||||
}
|
||||
}
|
||||
|
||||
/// Increment the counter by the given amount
|
||||
pub fn add(&self, delta: u64) {
|
||||
let _ = self
|
||||
.value
|
||||
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
|
||||
Some(current.saturating_add(delta))
|
||||
});
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is accessing this counter simultaneously.
|
||||
pub unsafe fn add(&self, delta: u64) {
|
||||
let ptr = self.value.get();
|
||||
unsafe {
|
||||
*ptr = (*ptr).saturating_add(delta);
|
||||
}
|
||||
}
|
||||
|
||||
/// Increment the counter by 1
|
||||
pub fn inc(&self) {
|
||||
self.add(1);
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is accessing this counter simultaneously.
|
||||
pub unsafe fn inc(&self) {
|
||||
unsafe {
|
||||
self.add(1);
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the current value of the counter
|
||||
pub fn get(&self) -> u64 {
|
||||
self.value.load(Ordering::Relaxed)
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is modifying this counter simultaneously.
|
||||
pub unsafe fn get(&self) -> u64 {
|
||||
let ptr = self.value.get();
|
||||
unsafe { *ptr }
|
||||
}
|
||||
|
||||
/// Reset the counter to zero
|
||||
pub fn reset(&self) {
|
||||
self.value.store(0, Ordering::Relaxed);
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is accessing this counter simultaneously.
|
||||
pub unsafe fn reset(&self) {
|
||||
let ptr = self.value.get();
|
||||
unsafe {
|
||||
*ptr = 0;
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the counter to a specific value
|
||||
pub fn set(&self, value: u64) {
|
||||
self.value.store(value, Ordering::Relaxed);
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is accessing this counter simultaneously.
|
||||
pub unsafe fn set(&self, value: u64) {
|
||||
let ptr = self.value.get();
|
||||
unsafe {
|
||||
*ptr = value;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// UnsafeCounter is Send + Sync because the safety is guaranteed by the caller
|
||||
unsafe impl Send for UnsafeCounter {}
|
||||
unsafe impl Sync for UnsafeCounter {}
|
||||
|
||||
/// MetricData contains both the counter and last update timestamp
|
||||
/// Uses UnsafeCell for lock-free access
|
||||
#[derive(Debug)]
|
||||
struct MetricData {
|
||||
counter: UnsafeCounter,
|
||||
last_updated: Mutex<Instant>,
|
||||
last_updated: UnsafeCell<Instant>,
|
||||
}
|
||||
|
||||
impl MetricData {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
counter: UnsafeCounter::new(),
|
||||
last_updated: Mutex::new(Instant::now()),
|
||||
last_updated: UnsafeCell::new(Instant::now()),
|
||||
}
|
||||
}
|
||||
|
||||
fn new_with_value(initial: u64) -> Self {
|
||||
Self {
|
||||
counter: UnsafeCounter::new_with_value(initial),
|
||||
last_updated: Mutex::new(Instant::now()),
|
||||
last_updated: UnsafeCell::new(Instant::now()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Update the last_updated timestamp
|
||||
fn touch(&self) {
|
||||
*self.last_updated.lock() = Instant::now();
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is accessing this timestamp simultaneously.
|
||||
unsafe fn touch(&self) {
|
||||
let ptr = self.last_updated.get();
|
||||
unsafe {
|
||||
*ptr = Instant::now();
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the last updated timestamp
|
||||
fn get_last_updated(&self) -> Instant {
|
||||
*self.last_updated.lock()
|
||||
/// # Safety
|
||||
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
|
||||
/// that no other thread is modifying this timestamp simultaneously.
|
||||
unsafe fn get_last_updated(&self) -> Instant {
|
||||
let ptr = self.last_updated.get();
|
||||
unsafe { *ptr }
|
||||
}
|
||||
}
|
||||
|
||||
// MetricData is Send + Sync because the safety is guaranteed by the caller
|
||||
unsafe impl Send for MetricData {}
|
||||
unsafe impl Sync for MetricData {}
|
||||
|
||||
/// MetricKey uniquely identifies a metric with its name and labels
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
struct MetricKey {
|
||||
@@ -508,31 +546,39 @@ impl CounterHandle {
|
||||
|
||||
/// Increment the counter by the given amount
|
||||
pub fn add(&self, delta: u64) {
|
||||
self.metric_data.counter.add(delta);
|
||||
self.metric_data.touch();
|
||||
unsafe {
|
||||
self.metric_data.counter.add(delta);
|
||||
self.metric_data.touch();
|
||||
}
|
||||
}
|
||||
|
||||
/// Increment the counter by 1
|
||||
pub fn inc(&self) {
|
||||
self.metric_data.counter.inc();
|
||||
self.metric_data.touch();
|
||||
unsafe {
|
||||
self.metric_data.counter.inc();
|
||||
self.metric_data.touch();
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the current value of the counter
|
||||
pub fn get(&self) -> u64 {
|
||||
self.metric_data.counter.get()
|
||||
unsafe { self.metric_data.counter.get() }
|
||||
}
|
||||
|
||||
/// Reset the counter to zero
|
||||
pub fn reset(&self) {
|
||||
self.metric_data.counter.reset();
|
||||
self.metric_data.touch();
|
||||
unsafe {
|
||||
self.metric_data.counter.reset();
|
||||
self.metric_data.touch();
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the counter to a specific value
|
||||
pub fn set(&self, value: u64) {
|
||||
self.metric_data.counter.set(value);
|
||||
self.metric_data.touch();
|
||||
unsafe {
|
||||
self.metric_data.counter.set(value);
|
||||
self.metric_data.touch();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -578,7 +624,7 @@ impl StatsManager {
|
||||
|
||||
counters.retain(|_, metric_data: &mut Arc<MetricData>| {
|
||||
Arc::strong_count(metric_data) > 1
|
||||
|| metric_data.get_last_updated() > cutoff_time
|
||||
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|
||||
});
|
||||
counters.shrink_to_fit();
|
||||
}
|
||||
@@ -616,7 +662,7 @@ impl StatsManager {
|
||||
let key = entry.key();
|
||||
let metric_data = entry.value();
|
||||
|
||||
let value = metric_data.counter.get();
|
||||
let value = unsafe { metric_data.counter.get() };
|
||||
|
||||
metrics.push(MetricSnapshot {
|
||||
name: key.name,
|
||||
@@ -649,7 +695,7 @@ impl StatsManager {
|
||||
let key = MetricKey::new(name, labels.clone());
|
||||
|
||||
if let Some(metric_data) = self.counters.get(&key) {
|
||||
let value = metric_data.counter.get();
|
||||
let value = unsafe { metric_data.counter.get() };
|
||||
Some(MetricSnapshot {
|
||||
name,
|
||||
labels: labels.clone(),
|
||||
@@ -750,15 +796,17 @@ mod tests {
|
||||
async fn test_unsafe_counter() {
|
||||
let counter = UnsafeCounter::new();
|
||||
|
||||
assert_eq!(counter.get(), 0);
|
||||
counter.inc();
|
||||
assert_eq!(counter.get(), 1);
|
||||
counter.add(5);
|
||||
assert_eq!(counter.get(), 6);
|
||||
counter.set(10);
|
||||
assert_eq!(counter.get(), 10);
|
||||
counter.reset();
|
||||
assert_eq!(counter.get(), 0);
|
||||
unsafe {
|
||||
assert_eq!(counter.get(), 0);
|
||||
counter.inc();
|
||||
assert_eq!(counter.get(), 1);
|
||||
counter.add(5);
|
||||
assert_eq!(counter.get(), 6);
|
||||
counter.set(10);
|
||||
assert_eq!(counter.get(), 10);
|
||||
counter.reset();
|
||||
assert_eq!(counter.get(), 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -903,7 +951,8 @@ mod tests {
|
||||
stats
|
||||
.counters
|
||||
.retain(|_, metric_data: &mut Arc<MetricData>| {
|
||||
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time
|
||||
Arc::strong_count(metric_data) > 1
|
||||
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|
||||
});
|
||||
|
||||
assert_eq!(stats.metric_count(), 1);
|
||||
@@ -913,33 +962,12 @@ mod tests {
|
||||
stats
|
||||
.counters
|
||||
.retain(|_, metric_data: &mut Arc<MetricData>| {
|
||||
Arc::strong_count(metric_data) > 1 || metric_data.get_last_updated() > cutoff_time
|
||||
Arc::strong_count(metric_data) > 1
|
||||
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|
||||
});
|
||||
assert_eq!(stats.metric_count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_counter_handle_concurrent_increment() {
|
||||
const THREADS: usize = 8;
|
||||
const INCREMENTS_PER_THREAD: usize = 10_000;
|
||||
|
||||
let stats = StatsManager::new();
|
||||
let counter = stats.get_simple_counter(MetricName::TrafficPacketsForwarded);
|
||||
|
||||
std::thread::scope(|scope| {
|
||||
for _ in 0..THREADS {
|
||||
let counter = counter.clone();
|
||||
scope.spawn(move || {
|
||||
for _ in 0..INCREMENTS_PER_THREAD {
|
||||
counter.inc();
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
assert_eq!(counter.get(), (THREADS * INCREMENTS_PER_THREAD) as u64);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stats_rpc_data_structures() {
|
||||
// Test GetStatsRequest
|
||||
|
||||
+27
-11
@@ -204,7 +204,7 @@ struct NetworkOptions {
|
||||
num_args = 0..=1,
|
||||
default_missing_value = "true"
|
||||
)]
|
||||
dhcp: Option<bool>,
|
||||
dhcp: Option<String>,
|
||||
|
||||
#[arg(
|
||||
short,
|
||||
@@ -907,8 +907,25 @@ impl NetworkOptions {
|
||||
cfg.set_network_identity(NetworkIdentity::new_credential(network_name));
|
||||
}
|
||||
|
||||
if let Some(dhcp) = self.dhcp {
|
||||
cfg.set_dhcp(dhcp);
|
||||
if let Some(ref dhcp) = self.dhcp {
|
||||
if dhcp == "true" || dhcp == "1" {
|
||||
cfg.set_dhcp(true);
|
||||
} else if dhcp == "false" || dhcp == "0" {
|
||||
cfg.set_dhcp(false);
|
||||
} else {
|
||||
// Treat as CIDR, e.g. "10.0.0.0/24"
|
||||
cfg.set_dhcp(true);
|
||||
let cidr: cidr::Ipv4Cidr = dhcp
|
||||
.parse()
|
||||
.with_context(|| format!("failed to parse dhcp cidr: {}", dhcp))?;
|
||||
if cidr.network_length() > 30 {
|
||||
anyhow::bail!(
|
||||
"dhcp cidr prefix length must be <= 30, got /{}",
|
||||
cidr.network_length()
|
||||
);
|
||||
}
|
||||
cfg.set_dhcp_cidr(Some(cidr));
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(ipv4) = &self.ipv4 {
|
||||
@@ -1614,7 +1631,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;
|
||||
@@ -1624,7 +1641,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;
|
||||
}
|
||||
|
||||
@@ -1644,13 +1661,12 @@ 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
|
||||
.context("failed to read config from stdin")?;
|
||||
TomlConfigLoader::new_from_str_with_source("stdin", stdin.as_str())?;
|
||||
_ = tokio::io::stdin().read_to_string(&mut stdin).await?;
|
||||
TomlConfigLoader::new_from_str(stdin.as_str())
|
||||
.with_context(|| "config source: stdin")?;
|
||||
} else {
|
||||
TomlConfigLoader::new(config_file)?;
|
||||
TomlConfigLoader::new(config_file)
|
||||
.with_context(|| format!("config source: {:?}", config_file))?;
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
@@ -829,7 +829,11 @@ impl Instance {
|
||||
let nic_ctx = self.nic_ctx.clone();
|
||||
let _peer_packet_receiver = self.peer_packet_receiver.clone();
|
||||
tokio::spawn(async move {
|
||||
let default_ipv4_addr = Ipv4Inet::new(Ipv4Addr::new(10, 126, 126, 0), 24).unwrap();
|
||||
let default_ipv4_addr = if let Some(dhcp_cidr) = global_ctx_c.config.get_dhcp_cidr() {
|
||||
Ipv4Inet::new(dhcp_cidr.first_address(), dhcp_cidr.network_length()).unwrap()
|
||||
} else {
|
||||
Ipv4Inet::new(Ipv4Addr::new(10, 126, 126, 0), 24).unwrap()
|
||||
};
|
||||
let mut current_dhcp_ip: Option<Ipv4Inet> = None;
|
||||
let mut next_sleep_time = 0;
|
||||
let nic_closed_notifier = Arc::new(Notify::new());
|
||||
@@ -864,7 +868,11 @@ impl Instance {
|
||||
used_ipv4.insert(peer_ipv4_addr.into());
|
||||
}
|
||||
|
||||
let dhcp_inet = used_ipv4.iter().next().unwrap_or(&default_ipv4_addr);
|
||||
let dhcp_inet = if global_ctx_c.config.get_dhcp_cidr().is_some() {
|
||||
&default_ipv4_addr
|
||||
} else {
|
||||
used_ipv4.iter().next().unwrap_or(&default_ipv4_addr)
|
||||
};
|
||||
// if old ip is already in this subnet and not conflicted, use it
|
||||
if let Some(ip) = current_dhcp_ip
|
||||
&& ip.network() == dhcp_inet.network()
|
||||
@@ -1810,4 +1818,83 @@ mod tests {
|
||||
|
||||
assert!(InstanceConfigPatcher::validate_public_ipv6_patch(&global_ctx, &patch).is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dhcp_cidr_allocates_ip_in_specified_subnet() {
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::common::config::{ConfigLoader, TomlConfigLoader};
|
||||
use crate::instance::instance::Instance;
|
||||
use crate::tunnel::common::tests::wait_for_condition;
|
||||
use crate::tunnel::ring::RingTunnelConnector;
|
||||
|
||||
// inst1: static IP, no DHCP (acts as a peer so DHCP on inst2 can proceed)
|
||||
let config1 = TomlConfigLoader::default();
|
||||
config1.set_inst_name("dhcp_test_inst1".to_owned());
|
||||
config1.set_ipv4(Some("192.168.200.1/24".parse().unwrap()));
|
||||
let mut flags1 = config1.get_flags();
|
||||
flags1.no_tun = true;
|
||||
config1.set_flags(flags1);
|
||||
config1.set_listeners(vec![]);
|
||||
|
||||
// inst2: DHCP enabled with specific CIDR
|
||||
let config2 = TomlConfigLoader::default();
|
||||
config2.set_inst_name("dhcp_test_inst2".to_owned());
|
||||
config2.set_dhcp(true);
|
||||
config2.set_dhcp_cidr(Some("172.20.0.0/24".parse().unwrap()));
|
||||
let mut flags2 = config2.get_flags();
|
||||
flags2.no_tun = true;
|
||||
config2.set_flags(flags2);
|
||||
config2.set_listeners(vec![]);
|
||||
|
||||
let mut inst1 = Instance::new(config1);
|
||||
let mut inst2 = Instance::new(config2);
|
||||
|
||||
inst1.run().await.unwrap();
|
||||
inst2.run().await.unwrap();
|
||||
|
||||
// Connect inst2 to inst1 via ring tunnel
|
||||
inst2
|
||||
.get_conn_manager()
|
||||
.add_connector(RingTunnelConnector::new(
|
||||
format!("ring://{}", inst1.id()).parse().unwrap(),
|
||||
));
|
||||
|
||||
// Wait for inst2 to see inst1 in routes
|
||||
let pm2 = inst2.get_peer_manager();
|
||||
wait_for_condition(
|
||||
|| async {
|
||||
let routes = pm2.list_routes().await;
|
||||
!routes.is_empty()
|
||||
},
|
||||
Duration::from_secs(5),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Wait for DHCP to allocate an IP on inst2
|
||||
let global_ctx2 = inst2.get_global_ctx();
|
||||
wait_for_condition(
|
||||
|| async { global_ctx2.get_ipv4().is_some() },
|
||||
Duration::from_secs(15),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Verify allocated IP is within the specified CIDR 172.20.0.0/24
|
||||
let allocated_ip = global_ctx2.get_ipv4().unwrap();
|
||||
let expected_cidr: cidr::Ipv4Cidr = "172.20.0.0/24".parse().unwrap();
|
||||
assert!(
|
||||
expected_cidr.contains(&allocated_ip.address()),
|
||||
"Allocated IP {:?} is not in expected CIDR {:?}",
|
||||
allocated_ip,
|
||||
expected_cidr
|
||||
);
|
||||
// Verify the network prefix length matches
|
||||
assert_eq!(
|
||||
allocated_ip.network_length(),
|
||||
expected_cidr.network_length(),
|
||||
"Allocated IP network length {} does not match expected {}",
|
||||
allocated_ip.network_length(),
|
||||
expected_cidr.network_length()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -637,6 +637,19 @@ impl NetworkConfig {
|
||||
);
|
||||
cfg.set_hostname(self.hostname.clone());
|
||||
cfg.set_dhcp(self.dhcp.unwrap_or_default());
|
||||
if let Some(ref dhcp_cidr) = self.dhcp_cidr {
|
||||
let cidr = dhcp_cidr
|
||||
.parse::<cidr::Ipv4Cidr>()
|
||||
.with_context(|| format!("failed to parse dhcp_cidr: {}", dhcp_cidr))?;
|
||||
if cidr.network_length() > 30 {
|
||||
anyhow::bail!(
|
||||
"dhcp_cidr prefix length must be <= 30, got /{}",
|
||||
cidr.network_length()
|
||||
);
|
||||
}
|
||||
cfg.set_dhcp(true);
|
||||
cfg.set_dhcp_cidr(Some(cidr));
|
||||
}
|
||||
cfg.set_inst_name(self.network_name.clone().unwrap_or_default());
|
||||
|
||||
// The web UI does not expose credential inputs directly, but imported/saved
|
||||
@@ -1019,6 +1032,7 @@ impl NetworkConfig {
|
||||
}
|
||||
|
||||
result.dhcp = Some(config.get_dhcp());
|
||||
result.dhcp_cidr = config.get_dhcp_cidr().map(|c| c.to_string());
|
||||
|
||||
let network_identity = config.get_network_identity();
|
||||
result.network_name = Some(network_identity.network_name.clone());
|
||||
|
||||
@@ -14,11 +14,11 @@ use std::{
|
||||
};
|
||||
|
||||
use dashmap::{DashMap, DashSet};
|
||||
use guarden::{Guard, defer};
|
||||
use guarden::defer;
|
||||
use tokio::{
|
||||
sync::{
|
||||
Mutex,
|
||||
mpsc::{self, Receiver, Sender, error::TrySendError},
|
||||
mpsc::{self, UnboundedReceiver, UnboundedSender},
|
||||
},
|
||||
task::JoinSet,
|
||||
};
|
||||
@@ -30,7 +30,7 @@ use crate::{
|
||||
error::Error,
|
||||
global_ctx::{ArcGlobalCtx, GlobalCtx, GlobalCtxEvent, NetworkIdentity, TrustedKeySource},
|
||||
join_joinset_background, shrink_dashmap,
|
||||
stats_manager::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager},
|
||||
stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
|
||||
token_bucket::TokenBucket,
|
||||
},
|
||||
peer_center::instance::{PeerCenterInstance, PeerMapWithPeerRpcManager},
|
||||
@@ -64,35 +64,6 @@ use super::{
|
||||
},
|
||||
};
|
||||
|
||||
const PEER_RPC_PACKET_QUEUE_CAPACITY: usize = 1024;
|
||||
|
||||
fn try_enqueue_peer_rpc_packet(
|
||||
sender: &Sender<ZCPacket>,
|
||||
packet: ZCPacket,
|
||||
dropped_packets: &CounterHandle,
|
||||
queue_name: &'static str,
|
||||
) -> bool {
|
||||
match sender.try_send(packet) {
|
||||
Ok(()) => true,
|
||||
Err(TrySendError::Full(_)) => {
|
||||
dropped_packets.inc();
|
||||
tracing::warn!(
|
||||
queue = queue_name,
|
||||
"drop peer rpc/control packet because queue is full"
|
||||
);
|
||||
false
|
||||
}
|
||||
Err(TrySendError::Closed(_)) => {
|
||||
dropped_packets.inc();
|
||||
tracing::warn!(
|
||||
queue = queue_name,
|
||||
"drop peer rpc/control packet because receiver is closed"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(&, Box, Arc)]
|
||||
pub trait GlobalForeignNetworkAccessor: Send + Sync + 'static {
|
||||
@@ -116,7 +87,7 @@ struct ForeignNetworkEntry {
|
||||
pm_packet_sender: Mutex<Option<PacketRecvChan>>,
|
||||
|
||||
peer_rpc: Arc<PeerRpcManager>,
|
||||
rpc_sender: Sender<ZCPacket>,
|
||||
rpc_sender: UnboundedSender<ZCPacket>,
|
||||
|
||||
packet_recv: Mutex<Option<PacketRecvChanReceiver>>,
|
||||
|
||||
@@ -341,12 +312,12 @@ impl ForeignNetworkEntry {
|
||||
fn build_rpc_tspt(
|
||||
my_peer_id: PeerId,
|
||||
peer_map: Arc<PeerMap>,
|
||||
) -> (Arc<PeerRpcManager>, Sender<ZCPacket>) {
|
||||
) -> (Arc<PeerRpcManager>, UnboundedSender<ZCPacket>) {
|
||||
struct RpcTransport {
|
||||
my_peer_id: PeerId,
|
||||
peer_map: Weak<PeerMap>,
|
||||
|
||||
packet_recv: Mutex<Receiver<ZCPacket>>,
|
||||
packet_recv: Mutex<UnboundedReceiver<ZCPacket>>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
@@ -388,8 +359,7 @@ impl ForeignNetworkEntry {
|
||||
}
|
||||
}
|
||||
|
||||
let (rpc_transport_sender, peer_rpc_tspt_recv) =
|
||||
mpsc::channel(PEER_RPC_PACKET_QUEUE_CAPACITY);
|
||||
let (rpc_transport_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel();
|
||||
let tspt = RpcTransport {
|
||||
my_peer_id,
|
||||
peer_map: Arc::downgrade(&peer_map),
|
||||
@@ -508,9 +478,6 @@ impl ForeignNetworkEntry {
|
||||
let rx_packets = self
|
||||
.stats_mgr
|
||||
.get_counter(MetricName::TrafficPacketsRx, label_set.clone());
|
||||
let rpc_queue_drops = self
|
||||
.stats_mgr
|
||||
.get_counter(MetricName::PeerRpcPacketQueueDrops, label_set.clone());
|
||||
|
||||
self.tasks.lock().await.spawn(async move {
|
||||
while let Ok(mut zc_packet) = recv_packet_from_chan(&mut recv).await {
|
||||
@@ -559,12 +526,7 @@ impl ForeignNetworkEntry {
|
||||
{
|
||||
rx_bytes.add(buf_len as u64);
|
||||
rx_packets.inc();
|
||||
try_enqueue_peer_rpc_packet(
|
||||
&rpc_sender,
|
||||
zc_packet,
|
||||
&rpc_queue_drops,
|
||||
"foreign_network_peer_rpc",
|
||||
);
|
||||
rpc_sender.send(zc_packet).unwrap();
|
||||
continue;
|
||||
}
|
||||
tracing::trace!(
|
||||
@@ -1274,7 +1236,7 @@ impl Drop for ForeignNetworkManager {
|
||||
pub mod tests {
|
||||
use crate::{
|
||||
common::global_ctx::tests::get_mock_global_ctx_with_network,
|
||||
common::stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
|
||||
common::stats_manager::{LabelSet, LabelType, MetricName},
|
||||
connector::udp_hole_punch::tests::{
|
||||
create_mock_peer_manager_with_mock_stun, replace_stun_info_collector,
|
||||
},
|
||||
@@ -1291,7 +1253,6 @@ pub mod tests {
|
||||
},
|
||||
};
|
||||
use std::{collections::HashMap, time::Duration};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use super::*;
|
||||
|
||||
@@ -1304,59 +1265,6 @@ pub mod tests {
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_enqueues_when_available() {
|
||||
let (sender, mut receiver) = mpsc::channel(1);
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"rpc"),
|
||||
&dropped_packets,
|
||||
"test_foreign_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 0);
|
||||
assert!(receiver.try_recv().is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_drops_when_full() {
|
||||
let (sender, _receiver) = mpsc::channel(1);
|
||||
sender
|
||||
.try_send(ZCPacket::new_with_payload(b"existing"))
|
||||
.unwrap();
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(!try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"overflow"),
|
||||
&dropped_packets,
|
||||
"test_foreign_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_drops_when_closed() {
|
||||
let (sender, receiver) = mpsc::channel(1);
|
||||
drop(receiver);
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(!try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"closed"),
|
||||
&dropped_packets,
|
||||
"test_foreign_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 1);
|
||||
}
|
||||
|
||||
async fn create_mock_peer_manager_for_foreign_network_ext(
|
||||
network: &str,
|
||||
secret: &str,
|
||||
|
||||
@@ -33,6 +33,7 @@ use super::{
|
||||
peer_session::{PeerSession, PeerSessionAction},
|
||||
traffic_metrics::AggregateTrafficMetrics,
|
||||
};
|
||||
use crate::utils::BoxExt;
|
||||
use crate::{
|
||||
common::{
|
||||
PeerId,
|
||||
@@ -379,9 +380,9 @@ impl PeerConn {
|
||||
session_filter,
|
||||
noise_handshake_result: None,
|
||||
|
||||
tunnel: Arc::new(Mutex::new(Box::new(
|
||||
guard!([mut mpsc_tunnel] mpsc_tunnel.close()),
|
||||
))),
|
||||
tunnel: Arc::new(Mutex::new(
|
||||
guard!([mut mpsc_tunnel] mpsc_tunnel.close()).boxed(),
|
||||
)),
|
||||
sink,
|
||||
recv: Mutex::new(Some(recv)),
|
||||
tunnel_info,
|
||||
|
||||
@@ -13,7 +13,7 @@ use std::{
|
||||
use tokio::{
|
||||
sync::{
|
||||
Mutex, RwLock,
|
||||
mpsc::{self, Receiver, Sender, error::TrySendError},
|
||||
mpsc::{self, UnboundedReceiver, UnboundedSender},
|
||||
},
|
||||
task::JoinSet,
|
||||
};
|
||||
@@ -72,43 +72,14 @@ use super::{
|
||||
route_trait::{ArcRoute, Route},
|
||||
};
|
||||
|
||||
const PEER_RPC_PACKET_QUEUE_CAPACITY: usize = 1024;
|
||||
|
||||
fn try_enqueue_peer_rpc_packet(
|
||||
sender: &Sender<ZCPacket>,
|
||||
packet: ZCPacket,
|
||||
dropped_packets: &CounterHandle,
|
||||
queue_name: &'static str,
|
||||
) -> bool {
|
||||
match sender.try_send(packet) {
|
||||
Ok(()) => true,
|
||||
Err(TrySendError::Full(_)) => {
|
||||
dropped_packets.inc();
|
||||
tracing::warn!(
|
||||
queue = queue_name,
|
||||
"drop peer rpc/control packet because queue is full"
|
||||
);
|
||||
false
|
||||
}
|
||||
Err(TrySendError::Closed(_)) => {
|
||||
dropped_packets.inc();
|
||||
tracing::warn!(
|
||||
queue = queue_name,
|
||||
"drop peer rpc/control packet because receiver is closed"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct RpcTransport {
|
||||
my_peer_id: PeerId,
|
||||
peers: Weak<PeerMap>,
|
||||
// TODO: this seems can be removed
|
||||
foreign_peers: Mutex<Option<Weak<ForeignNetworkClient>>>,
|
||||
|
||||
packet_recv: Mutex<Receiver<ZCPacket>>,
|
||||
peer_rpc_tspt_sender: Sender<ZCPacket>,
|
||||
packet_recv: Mutex<UnboundedReceiver<ZCPacket>>,
|
||||
peer_rpc_tspt_sender: UnboundedSender<ZCPacket>,
|
||||
|
||||
encryptor: Arc<dyn Encryptor>,
|
||||
is_secure_mode_enabled: bool,
|
||||
@@ -302,8 +273,7 @@ impl PeerManager {
|
||||
.unwrap_or(false);
|
||||
|
||||
// TODO: remove these because we have impl pipeline processor.
|
||||
let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) =
|
||||
mpsc::channel(PEER_RPC_PACKET_QUEUE_CAPACITY);
|
||||
let (peer_rpc_tspt_sender, peer_rpc_tspt_recv) = mpsc::unbounded_channel();
|
||||
let rpc_tspt = Arc::new(RpcTransport {
|
||||
my_peer_id,
|
||||
peers: Arc::downgrade(&peers),
|
||||
@@ -1273,8 +1243,7 @@ impl PeerManager {
|
||||
|
||||
// for peer rpc packet
|
||||
struct PeerRpcPacketProcessor {
|
||||
peer_rpc_tspt_sender: Sender<ZCPacket>,
|
||||
dropped_packets: CounterHandle,
|
||||
peer_rpc_tspt_sender: UnboundedSender<ZCPacket>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
@@ -1285,27 +1254,15 @@ impl PeerManager {
|
||||
|| hdr.packet_type == PacketType::RpcReq as u8
|
||||
|| hdr.packet_type == PacketType::RpcResp as u8
|
||||
{
|
||||
try_enqueue_peer_rpc_packet(
|
||||
&self.peer_rpc_tspt_sender,
|
||||
packet,
|
||||
&self.dropped_packets,
|
||||
"local_peer_rpc",
|
||||
);
|
||||
self.peer_rpc_tspt_sender.send(packet).unwrap();
|
||||
None
|
||||
} else {
|
||||
Some(packet)
|
||||
}
|
||||
}
|
||||
}
|
||||
let peer_rpc_queue_drops = self.global_ctx.stats_manager().get_counter(
|
||||
MetricName::PeerRpcPacketQueueDrops,
|
||||
LabelSet::new().with_label_type(LabelType::NetworkName(
|
||||
self.global_ctx.get_network_name().to_string(),
|
||||
)),
|
||||
);
|
||||
self.add_packet_process_pipeline(Box::new(PeerRpcPacketProcessor {
|
||||
peer_rpc_tspt_sender: self.peer_rpc_tspt.peer_rpc_tspt_sender.clone(),
|
||||
dropped_packets: peer_rpc_queue_drops,
|
||||
}))
|
||||
.await;
|
||||
}
|
||||
@@ -1576,22 +1533,9 @@ 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 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) {
|
||||
let send_result = 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
|
||||
@@ -2241,7 +2185,6 @@ impl PeerManager {
|
||||
mod tests {
|
||||
use base64::Engine;
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
fmt::Debug,
|
||||
sync::Arc,
|
||||
time::{Duration, Instant},
|
||||
@@ -2249,10 +2192,9 @@ mod tests {
|
||||
|
||||
use crate::{
|
||||
common::{
|
||||
PeerId,
|
||||
config::Flags,
|
||||
global_ctx::{NetworkIdentity, tests::get_mock_global_ctx},
|
||||
stats_manager::{LabelSet, LabelType, MetricName, StatsManager},
|
||||
stats_manager::{LabelSet, LabelType, MetricName},
|
||||
},
|
||||
connector::{
|
||||
create_connector_by_url, direct::PeerManagerForDirectConnector,
|
||||
@@ -2264,7 +2206,7 @@ mod tests {
|
||||
peer_conn::tests::set_secure_mode_cfg,
|
||||
peer_manager::RouteAlgoType,
|
||||
peer_rpc::tests::register_service,
|
||||
route_trait::{NextHopPolicy, RouteCostCalculatorInterface},
|
||||
route_trait::NextHopPolicy,
|
||||
tests::{
|
||||
connect_peer_manager, create_mock_peer_manager_with_name, wait_route_appear,
|
||||
wait_route_appear_with_cost,
|
||||
@@ -2283,9 +2225,7 @@ mod tests {
|
||||
},
|
||||
};
|
||||
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use super::{PeerManager, try_enqueue_peer_rpc_packet};
|
||||
use super::PeerManager;
|
||||
|
||||
async fn create_lazy_peer_manager() -> Arc<PeerManager> {
|
||||
let peer_mgr = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
|
||||
@@ -2310,69 +2250,6 @@ mod tests {
|
||||
))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_enqueues_when_available() {
|
||||
let (sender, mut receiver) = mpsc::channel(1);
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"rpc"),
|
||||
&dropped_packets,
|
||||
"test_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 0);
|
||||
assert!(receiver.try_recv().is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_drops_when_full() {
|
||||
let (sender, _receiver) = mpsc::channel(1);
|
||||
sender
|
||||
.try_send(ZCPacket::new_with_payload(b"existing"))
|
||||
.unwrap();
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(!try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"overflow"),
|
||||
&dropped_packets,
|
||||
"test_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_rpc_queue_helper_drops_when_closed() {
|
||||
let (sender, receiver) = mpsc::channel(1);
|
||||
drop(receiver);
|
||||
let stats_manager = StatsManager::new();
|
||||
let dropped_packets = stats_manager.get_simple_counter(MetricName::PeerRpcPacketQueueDrops);
|
||||
|
||||
assert!(!try_enqueue_peer_rpc_packet(
|
||||
&sender,
|
||||
ZCPacket::new_with_payload(b"closed"),
|
||||
&dropped_packets,
|
||||
"test_peer_rpc",
|
||||
));
|
||||
|
||||
assert_eq!(dropped_packets.get(), 1);
|
||||
}
|
||||
|
||||
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));
|
||||
@@ -2780,109 +2657,6 @@ 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;
|
||||
|
||||
@@ -102,6 +102,7 @@ message NetworkConfig {
|
||||
optional bool disable_relay_data = 65;
|
||||
optional bool enable_udp_broadcast_relay = 66;
|
||||
optional uint32 socket_mark = 67;
|
||||
optional string dhcp_cidr = 68;
|
||||
}
|
||||
|
||||
message PortForwardConfig {
|
||||
|
||||
@@ -12,7 +12,7 @@ use std::{
|
||||
sync::Arc,
|
||||
task::{Context as TaskContext, Poll},
|
||||
};
|
||||
use tokio::{io::AsyncReadExt, net::TcpStream};
|
||||
use tokio::{io::AsyncReadExt, net::TcpStream, sync::Mutex};
|
||||
|
||||
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<stack::Stack>>,
|
||||
stack_map: DashMap<String, Arc<Mutex<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<stack::Stack>, TunnelError> {
|
||||
) -> Result<Arc<Mutex<stack::Stack>>, TunnelError> {
|
||||
let local_socket_addr = accept_result.local_addr;
|
||||
|
||||
let interface_name = &accept_result.interface_name;
|
||||
@@ -158,38 +158,29 @@ impl FakeTcpTunnelListener {
|
||||
IpAddr::V6(ip) => (None, Some(ip)),
|
||||
};
|
||||
|
||||
if let Some(entry) = self.stack_map.get(interface_name) {
|
||||
let stack = entry.clone();
|
||||
drop(entry);
|
||||
|
||||
if !stack.is_closed() {
|
||||
return Ok(stack);
|
||||
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
|
||||
}
|
||||
};
|
||||
|
||||
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)
|
||||
Ok(ret)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -224,29 +215,19 @@ 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, 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);
|
||||
};
|
||||
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;
|
||||
|
||||
tracing::info!(
|
||||
?res,
|
||||
@@ -255,7 +236,7 @@ impl TunnelListener for FakeTcpTunnelListener {
|
||||
);
|
||||
|
||||
let info = TunnelInfo {
|
||||
tunnel_type: get_faketcp_tunnel_type_str(stack.driver_type()),
|
||||
tunnel_type: get_faketcp_tunnel_type_str(stack.lock().await.driver_type()),
|
||||
local_addr: Some(self.local_url().into()),
|
||||
remote_addr: Some(
|
||||
crate::tunnel::build_url_from_socket_addr(
|
||||
@@ -373,14 +354,12 @@ 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 stack = stack::Stack::new(tun, local_ip, local_ip6, mac);
|
||||
let mut stack = stack::Stack::new(tun, local_ip, local_ip6, mac);
|
||||
let driver_type = stack.driver_type();
|
||||
|
||||
let socket = stack
|
||||
.try_alloc_established_socket(local_addr, remote_addr, stack::State::SynSent)
|
||||
.ok_or(TunnelError::InternalError(
|
||||
"FakeTCP stack closed while allocating socket".into(),
|
||||
))?;
|
||||
.alloc_established_socket(local_addr, remote_addr, stack::State::SynSent)
|
||||
.await;
|
||||
|
||||
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::{error, info, trace, warn};
|
||||
use tracing::{info, trace, warn};
|
||||
|
||||
const TIMEOUT: time::Duration = time::Duration::from_secs(1);
|
||||
const RETRIES: usize = 6;
|
||||
@@ -83,33 +83,13 @@ impl AddrTuple {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct StackState {
|
||||
tuples: HashMap<AddrTuple, flume::Sender<Bytes>>,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
struct Shared {
|
||||
state: RwLock<StackState>,
|
||||
tuples: RwLock<HashMap<AddrTuple, flume::Sender<Bytes>>>,
|
||||
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,
|
||||
@@ -373,17 +353,7 @@ impl Drop for Socket {
|
||||
fn drop(&mut self) {
|
||||
let tuple = AddrTuple::new(self.local_addr, self.remote_addr);
|
||||
// dissociates ourself from the dispatch map
|
||||
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");
|
||||
}
|
||||
}
|
||||
assert!(self.shared.tuples.write().unwrap().remove(&tuple).is_some());
|
||||
// purge cache
|
||||
let _ = self.shared.tuples_purge.send(tuple);
|
||||
|
||||
@@ -430,7 +400,7 @@ impl Stack {
|
||||
) -> Stack {
|
||||
let (tuples_purge_tx, _tuples_purge_rx) = broadcast::channel(16);
|
||||
let shared = Arc::new(Shared {
|
||||
state: RwLock::new(StackState::default()),
|
||||
tuples: RwLock::new(HashMap::new()),
|
||||
tun: tun.clone(),
|
||||
listening: RwLock::new(HashSet::new()),
|
||||
tuples_purge: tuples_purge_tx.clone(),
|
||||
@@ -456,31 +426,19 @@ 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 fn try_alloc_established_socket(
|
||||
&self,
|
||||
pub async fn alloc_established_socket(
|
||||
&mut self,
|
||||
local_addr: SocketAddr,
|
||||
remote_addr: SocketAddr,
|
||||
state: State,
|
||||
) -> Option<Socket> {
|
||||
) -> Socket {
|
||||
let tuple = AddrTuple::new(local_addr, remote_addr);
|
||||
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 mut tuples = self.shared.tuples.write().unwrap();
|
||||
let (sock, incoming) = Socket::new(
|
||||
self.shared.clone(),
|
||||
// self.shared.tun.choose(&mut rng).unwrap().clone(),
|
||||
@@ -492,8 +450,8 @@ impl Stack {
|
||||
Some(0), // Initial ACK
|
||||
state,
|
||||
);
|
||||
assert!(stack_state.tuples.insert(tuple, incoming).is_none());
|
||||
Some(sock)
|
||||
assert!(tuples.insert(tuple, incoming).is_none());
|
||||
sock
|
||||
}
|
||||
|
||||
async fn reader_task(
|
||||
@@ -508,22 +466,7 @@ impl Stack {
|
||||
|
||||
tokio::select! {
|
||||
size = tun.recv(&mut buf) => {
|
||||
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;
|
||||
}
|
||||
};
|
||||
let size = size.unwrap();
|
||||
tracing::trace!(len = size, ?buf, "PnetTun received packet");
|
||||
let buf = buf.split().freeze();
|
||||
|
||||
@@ -551,8 +494,8 @@ impl Stack {
|
||||
} else {
|
||||
trace!("Cache miss, checking the shared tuples table for connection");
|
||||
let sender = {
|
||||
let state = shared.state.read().unwrap();
|
||||
state.tuples.get(&tuple).cloned()
|
||||
let tuples = shared.tuples.read().unwrap();
|
||||
tuples.get(&tuple).cloned()
|
||||
};
|
||||
|
||||
if let Some(c) = sender {
|
||||
@@ -589,107 +532,11 @@ impl Stack {
|
||||
}
|
||||
},
|
||||
tuple = tuples_purge.recv() => {
|
||||
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;
|
||||
}
|
||||
}
|
||||
let tuple = tuple.unwrap();
|
||||
tuples.remove(&tuple);
|
||||
trace!("Removed cached tuple: {:?}", tuple);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
|
||||
+2
-246
@@ -23,250 +23,6 @@ use std::{net::SocketAddr, sync::Arc, time::Duration};
|
||||
use tokio::net::UdpSocket;
|
||||
|
||||
// region config
|
||||
mod crypto {
|
||||
use crate::utils::BoxExt;
|
||||
use bytes::{Buf, BytesMut};
|
||||
use quinn_proto::crypto::{
|
||||
ClientConfig, ExportKeyingMaterialError, KeyPair, Keys, ServerConfig, Session,
|
||||
UnsupportedVersion,
|
||||
};
|
||||
use quinn_proto::transport_parameters::TransportParameters;
|
||||
use quinn_proto::{
|
||||
ConnectError, ConnectionId, Side, TransportError,
|
||||
crypto::{CryptoError, HeaderKey, PacketKey},
|
||||
};
|
||||
use seahash::SeaHasher;
|
||||
use std::any::Any;
|
||||
use std::{hash::Hasher, sync::Arc};
|
||||
use tracing::{error, instrument, trace};
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct CryptoKey;
|
||||
|
||||
impl CryptoKey {
|
||||
fn header(self) -> KeyPair<Box<dyn HeaderKey>> {
|
||||
KeyPair {
|
||||
local: Box::new(self),
|
||||
remote: Box::new(self),
|
||||
}
|
||||
}
|
||||
|
||||
fn packet(self) -> KeyPair<Box<dyn PacketKey>> {
|
||||
KeyPair {
|
||||
local: Box::new(self),
|
||||
remote: Box::new(self),
|
||||
}
|
||||
}
|
||||
|
||||
fn keys(self) -> Keys {
|
||||
Keys {
|
||||
header: self.header(),
|
||||
packet: self.packet(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl HeaderKey for CryptoKey {
|
||||
fn decrypt(&self, _: usize, _: &mut [u8]) {}
|
||||
fn encrypt(&self, _: usize, _: &mut [u8]) {}
|
||||
fn sample_size(&self) -> usize {
|
||||
0
|
||||
}
|
||||
}
|
||||
|
||||
impl CryptoKey {
|
||||
fn checksum(slices: &[&[u8]]) -> u64 {
|
||||
let mut hasher = SeaHasher::default();
|
||||
for slice in slices {
|
||||
hasher.write(&(slice.len() as u64).to_le_bytes());
|
||||
hasher.write(slice);
|
||||
}
|
||||
hasher.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PacketKey for CryptoKey {
|
||||
#[instrument(level = "trace")]
|
||||
fn encrypt(&self, packet: u64, buf: &mut [u8], header_len: usize) {
|
||||
let (header, rest) = buf.split_at_mut(header_len);
|
||||
let (payload, tag) = rest.split_at_mut(rest.len() - self.tag_len());
|
||||
let checksum = Self::checksum(&[header, payload]);
|
||||
tag.copy_from_slice(&checksum.to_be_bytes());
|
||||
trace!(checksum, ?header, ?payload, ?tag);
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn decrypt(
|
||||
&self,
|
||||
packet: u64,
|
||||
header: &[u8],
|
||||
payload: &mut BytesMut,
|
||||
) -> Result<(), CryptoError> {
|
||||
let tag = payload.split_off(payload.len() - self.tag_len()).get_u64();
|
||||
trace!(tag, ?payload);
|
||||
let checksum = Self::checksum(&[header, payload]);
|
||||
if checksum != tag {
|
||||
error!(tag, checksum, "checksum mismatch");
|
||||
return Err(CryptoError);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn tag_len(&self) -> usize {
|
||||
8
|
||||
}
|
||||
|
||||
fn confidentiality_limit(&self) -> u64 {
|
||||
u64::MAX
|
||||
}
|
||||
|
||||
fn integrity_limit(&self) -> u64 {
|
||||
1 << 36
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum HandshakeState {
|
||||
EmitInitial,
|
||||
EmitHandshake,
|
||||
Done,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct QuicSession {
|
||||
side: Side,
|
||||
state: HandshakeState,
|
||||
local: TransportParameters,
|
||||
remote: Option<TransportParameters>,
|
||||
}
|
||||
|
||||
impl QuicSession {
|
||||
fn new(side: Side, params: TransportParameters) -> Self {
|
||||
Self {
|
||||
side,
|
||||
state: HandshakeState::EmitInitial,
|
||||
local: params,
|
||||
remote: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Session for QuicSession {
|
||||
fn initial_keys(&self, _: &ConnectionId, _: Side) -> Keys {
|
||||
CryptoKey.keys()
|
||||
}
|
||||
|
||||
fn handshake_data(&self) -> Option<Box<dyn Any>> {
|
||||
self.remote.map(|params| params.boxed() as _)
|
||||
}
|
||||
|
||||
fn peer_identity(&self) -> Option<Box<dyn Any>> {
|
||||
None
|
||||
}
|
||||
|
||||
fn early_crypto(&self) -> Option<(Box<dyn HeaderKey>, Box<dyn PacketKey>)> {
|
||||
None
|
||||
}
|
||||
|
||||
fn early_data_accepted(&self) -> Option<bool> {
|
||||
Some(false)
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn is_handshaking(&self) -> bool {
|
||||
self.remote.is_none() || self.state != HandshakeState::Done
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn read_handshake(&mut self, mut buf: &[u8]) -> Result<bool, TransportError> {
|
||||
if self.remote.is_none() {
|
||||
self.remote = Some(
|
||||
TransportParameters::read(self.side, &mut buf)
|
||||
.expect("failed to read transport parameters"),
|
||||
);
|
||||
}
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn transport_parameters(&self) -> Result<Option<TransportParameters>, TransportError> {
|
||||
Ok(self.remote)
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn write_handshake(&mut self, buf: &mut Vec<u8>) -> Option<Keys> {
|
||||
match self.state {
|
||||
HandshakeState::EmitInitial => {
|
||||
if self.side.is_client() {
|
||||
self.local.write(buf);
|
||||
}
|
||||
self.state = HandshakeState::EmitHandshake;
|
||||
Some(CryptoKey.keys())
|
||||
}
|
||||
HandshakeState::EmitHandshake => {
|
||||
if self.side.is_server() {
|
||||
self.local.write(buf);
|
||||
}
|
||||
self.state = HandshakeState::Done;
|
||||
Some(CryptoKey.keys())
|
||||
}
|
||||
HandshakeState::Done => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn next_1rtt_keys(&mut self) -> Option<KeyPair<Box<dyn PacketKey>>> {
|
||||
Some(CryptoKey.packet())
|
||||
}
|
||||
|
||||
fn is_valid_retry(&self, _: &ConnectionId, _: &[u8], _: &[u8]) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn export_keying_material(
|
||||
&self,
|
||||
_: &mut [u8],
|
||||
_: &[u8],
|
||||
_: &[u8],
|
||||
) -> Result<(), ExportKeyingMaterialError> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct CryptoConfig;
|
||||
|
||||
impl ClientConfig for CryptoConfig {
|
||||
#[instrument(level = "trace")]
|
||||
fn start_session(
|
||||
self: Arc<Self>,
|
||||
version: u32,
|
||||
server_name: &str,
|
||||
params: &TransportParameters,
|
||||
) -> Result<Box<dyn Session>, ConnectError> {
|
||||
Ok(Box::new(QuicSession::new(Side::Client, *params)))
|
||||
}
|
||||
}
|
||||
|
||||
impl ServerConfig for CryptoConfig {
|
||||
fn initial_keys(&self, _: u32, _: &ConnectionId) -> Result<Keys, UnsupportedVersion> {
|
||||
Ok(CryptoKey.keys())
|
||||
}
|
||||
|
||||
fn retry_tag(&self, _: u32, _: &ConnectionId, _: &[u8]) -> [u8; 16] {
|
||||
[0u8; 16]
|
||||
}
|
||||
|
||||
#[instrument(level = "trace")]
|
||||
fn start_session(
|
||||
self: Arc<Self>,
|
||||
version: u32,
|
||||
params: &TransportParameters,
|
||||
) -> Box<dyn Session> {
|
||||
Box::new(QuicSession::new(Side::Server, *params))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn transport_config() -> Arc<TransportConfig> {
|
||||
let mut config = TransportConfig::default();
|
||||
|
||||
@@ -283,13 +39,13 @@ pub fn transport_config() -> Arc<TransportConfig> {
|
||||
}
|
||||
|
||||
pub fn server_config() -> ServerConfig {
|
||||
let mut config = ServerConfig::with_crypto(Arc::new(crypto::CryptoConfig));
|
||||
let mut config = quinn_plaintext::server_config();
|
||||
config.transport_config(transport_config());
|
||||
config
|
||||
}
|
||||
|
||||
pub fn client_config() -> ClientConfig {
|
||||
let mut config = ClientConfig::new(Arc::new(crypto::CryptoConfig));
|
||||
let mut config = quinn_plaintext::client_config();
|
||||
config.transport_config(transport_config());
|
||||
config
|
||||
}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering::Relaxed};
|
||||
use std::{
|
||||
cell::UnsafeCell,
|
||||
sync::atomic::{AtomicU32, Ordering::Relaxed},
|
||||
};
|
||||
|
||||
pub struct WindowLatency {
|
||||
latency_us_window: Vec<AtomicU32>,
|
||||
@@ -60,30 +63,34 @@ impl WindowLatency {
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct Throughput {
|
||||
tx_bytes: AtomicU64,
|
||||
rx_bytes: AtomicU64,
|
||||
tx_packets: AtomicU64,
|
||||
rx_packets: AtomicU64,
|
||||
tx_bytes: UnsafeCell<u64>,
|
||||
rx_bytes: UnsafeCell<u64>,
|
||||
tx_packets: UnsafeCell<u64>,
|
||||
rx_packets: UnsafeCell<u64>,
|
||||
}
|
||||
|
||||
impl Clone for Throughput {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
tx_bytes: AtomicU64::new(self.tx_bytes()),
|
||||
rx_bytes: AtomicU64::new(self.rx_bytes()),
|
||||
tx_packets: AtomicU64::new(self.tx_packets()),
|
||||
rx_packets: AtomicU64::new(self.rx_packets()),
|
||||
tx_bytes: UnsafeCell::new(unsafe { *self.tx_bytes.get() }),
|
||||
rx_bytes: UnsafeCell::new(unsafe { *self.rx_bytes.get() }),
|
||||
tx_packets: UnsafeCell::new(unsafe { *self.tx_packets.get() }),
|
||||
rx_packets: UnsafeCell::new(unsafe { *self.rx_packets.get() }),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// add sync::Send and sync::Sync traits to Throughput
|
||||
unsafe impl Send for Throughput {}
|
||||
unsafe impl Sync for Throughput {}
|
||||
|
||||
impl Default for Throughput {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
tx_bytes: AtomicU64::new(0),
|
||||
rx_bytes: AtomicU64::new(0),
|
||||
tx_packets: AtomicU64::new(0),
|
||||
rx_packets: AtomicU64::new(0),
|
||||
tx_bytes: UnsafeCell::new(0),
|
||||
rx_bytes: UnsafeCell::new(0),
|
||||
tx_packets: UnsafeCell::new(0),
|
||||
rx_packets: UnsafeCell::new(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -94,68 +101,32 @@ impl Throughput {
|
||||
}
|
||||
|
||||
pub fn tx_bytes(&self) -> u64 {
|
||||
self.tx_bytes.load(Relaxed)
|
||||
unsafe { *self.tx_bytes.get() }
|
||||
}
|
||||
|
||||
pub fn rx_bytes(&self) -> u64 {
|
||||
self.rx_bytes.load(Relaxed)
|
||||
unsafe { *self.rx_bytes.get() }
|
||||
}
|
||||
|
||||
pub fn tx_packets(&self) -> u64 {
|
||||
self.tx_packets.load(Relaxed)
|
||||
unsafe { *self.tx_packets.get() }
|
||||
}
|
||||
|
||||
pub fn rx_packets(&self) -> u64 {
|
||||
self.rx_packets.load(Relaxed)
|
||||
unsafe { *self.rx_packets.get() }
|
||||
}
|
||||
|
||||
pub fn record_tx_bytes(&self, bytes: u64) {
|
||||
self.tx_bytes.fetch_add(bytes, Relaxed);
|
||||
self.tx_packets.fetch_add(1, Relaxed);
|
||||
unsafe {
|
||||
*self.tx_bytes.get() += bytes;
|
||||
*self.tx_packets.get() += 1;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn record_rx_bytes(&self, bytes: u64) {
|
||||
self.rx_bytes.fetch_add(bytes, Relaxed);
|
||||
self.rx_packets.fetch_add(1, Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Throughput;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[test]
|
||||
fn throughput_records_concurrent_tx_and_rx() {
|
||||
const THREADS: usize = 8;
|
||||
const RECORDS_PER_THREAD: usize = 10_000;
|
||||
const TX_BYTES_PER_RECORD: u64 = 3;
|
||||
const RX_BYTES_PER_RECORD: u64 = 7;
|
||||
|
||||
let throughput = Arc::new(Throughput::new());
|
||||
|
||||
std::thread::scope(|scope| {
|
||||
for _ in 0..THREADS {
|
||||
let throughput = Arc::clone(&throughput);
|
||||
scope.spawn(move || {
|
||||
for _ in 0..RECORDS_PER_THREAD {
|
||||
throughput.record_tx_bytes(TX_BYTES_PER_RECORD);
|
||||
throughput.record_rx_bytes(RX_BYTES_PER_RECORD);
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
|
||||
let expected_packets = (THREADS * RECORDS_PER_THREAD) as u64;
|
||||
assert_eq!(throughput.tx_packets(), expected_packets);
|
||||
assert_eq!(throughput.rx_packets(), expected_packets);
|
||||
assert_eq!(
|
||||
throughput.tx_bytes(),
|
||||
expected_packets * TX_BYTES_PER_RECORD
|
||||
);
|
||||
assert_eq!(
|
||||
throughput.rx_bytes(),
|
||||
expected_packets * RX_BYTES_PER_RECORD
|
||||
);
|
||||
unsafe {
|
||||
*self.rx_bytes.get() += bytes;
|
||||
*self.rx_packets.get() += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -118,7 +118,6 @@ impl WsTunnelListener {
|
||||
|
||||
let (request, stream) = ServerBuilder::new()
|
||||
.limits(Limits::unlimited())
|
||||
.max_headers(128)
|
||||
.accept(stream)
|
||||
.await?;
|
||||
|
||||
@@ -253,8 +252,7 @@ impl WsTunnelConnector {
|
||||
),
|
||||
};
|
||||
|
||||
let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap())
|
||||
.max_headers(128);
|
||||
let c = ClientBuilder::from_uri(http::Uri::try_from(addr.to_string()).unwrap());
|
||||
let stream: MaybeTlsStream<TcpStream> = if is_wss {
|
||||
init_crypto_provider();
|
||||
let tls_conn =
|
||||
|
||||
Reference in New Issue
Block a user