Compare commits

..
Author SHA1 Message Date
copilot-swe-agent[bot]andGitHub 003fdefc63 fix: address review comments - CIDR validation, dhcp_cidr implies DHCP, fmt fixes 2026-06-11 16:54:40 +00:00
copilot-swe-agent[bot]andGitHub a25c411249 test: add test verifying DHCP allocates IP within specified CIDR subnet 2026-06-08 15:54:22 +00:00
copilot-swe-agent[bot]andGitHub 0c11cefc04 fix: return error instead of silently ignoring invalid dhcp_cidr in launcher 2026-06-08 15:35:30 +00:00
copilot-swe-agent[bot]andGitHub 8181830902 feat: add DHCP subnet (CIDR) configuration parameter
Allow -d to optionally accept a CIDR subnet (e.g. -d 10.0.0.0/24)
to fix the DHCP address range. When specified, the DHCP allocator
uses the configured subnet instead of deriving it from peer IPs.

Changes:
- Add dhcp_cidr field to config struct and ConfigLoader trait
- Modify -d CLI arg to accept optional CIDR string value
- Update DHCP logic to use configured CIDR as default subnet
- Add dhcp_cidr to NetworkConfig proto message
- Update launcher to handle dhcp_cidr in gen_config/new_from_config
- Update i18n translations
2026-06-08 15:31:53 +00:00
copilot-swe-agent[bot]andGitHub 0a8c95879b Initial plan 2026-06-08 15:22:59 +00:00
120 changed files with 2185 additions and 17483 deletions
+6 -13
View File
@@ -33,19 +33,6 @@ runs:
sudo apt-get install -qqy build-essential mold musl-tools
shell: bash
- name: Setup protoc
uses: arduino/setup-protoc@v3
with:
version: '35.1'
# GitHub repo token to use to avoid rate limiter
repo-token: ${{ inputs.token }}
- name: Verify protoc version
run: |
version="$(protoc --version | tr -d '\r')"
test "$version" = "libprotoc 35.1"
shell: bash
- name: Setup Frontend Environment
if: ${{ inputs.pnpm == 'true' }}
uses: ./.github/actions/prepare-pnpm
@@ -95,3 +82,9 @@ runs:
ar x libgcc.a _ctzsi2.o _clz.o _bswapsi2.o
ar rcs libctz.a _ctzsi2.o _clz.o _bswapsi2.o
shell: bash
- name: Setup protoc
uses: arduino/setup-protoc@v3
with:
# GitHub repo token to use to avoid rate limiter
repo-token: ${{ inputs.token }}
+2 -2
View File
@@ -41,8 +41,8 @@ runs:
pnpm -r install
if [ -n "${{ inputs.build-filter }}" ]; then
echo "Building with filter: ${{ inputs.build-filter }}"
pnpm -r --workspace-concurrency=1 --filter "${{ inputs.build-filter }}" build
pnpm -r --filter "${{ inputs.build-filter }}" build
else
echo "No build filter provided, building all packages"
pnpm -r --workspace-concurrency=1 build
pnpm -r build
fi
+1 -1
View File
@@ -35,7 +35,6 @@ jobs:
with:
gui: false
pnpm: false
token: ${{ secrets.GITHUB_TOKEN }}
- uses: actions-rust-lang/setup-rust-toolchain@v1
with:
@@ -244,3 +243,4 @@ jobs:
ohpm publish easytier-release.har
fi
curl --header "Content-Type: application/json" --request POST --data "{}" ${{ secrets.CODEARTS_WEBHOOKS }}
-1
View File
@@ -34,7 +34,6 @@ easytier-panic.log
# web
node_modules
easytier-web/frontend-lib/src/generated/
.vite
Generated
+26 -231
View File
@@ -129,12 +129,6 @@ dependencies = [
"libc",
]
[[package]]
name = "anes"
version = "0.1.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299"
[[package]]
name = "anstream"
version = "0.6.15"
@@ -247,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"
@@ -931,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",
@@ -1145,12 +1129,6 @@ dependencies = [
"toml 0.9.12+spec-1.1.0",
]
[[package]]
name = "cast"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
[[package]]
name = "cc"
version = "1.2.10"
@@ -1260,33 +1238,6 @@ dependencies = [
"windows-targets 0.52.6",
]
[[package]]
name = "ciborium"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e"
dependencies = [
"ciborium-io",
"ciborium-ll",
"serde",
]
[[package]]
name = "ciborium-io"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757"
[[package]]
name = "ciborium-ll"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9"
dependencies = [
"ciborium-io",
"half",
]
[[package]]
name = "cidr"
version = "0.3.1"
@@ -1632,42 +1583,6 @@ dependencies = [
"cfg-if",
]
[[package]]
name = "criterion"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f"
dependencies = [
"anes",
"cast",
"ciborium",
"clap",
"criterion-plot",
"is-terminal",
"itertools 0.10.5",
"num-traits",
"once_cell",
"oorandom",
"plotters",
"rayon",
"regex",
"serde",
"serde_derive",
"serde_json",
"tinytemplate",
"walkdir",
]
[[package]]
name = "criterion-plot"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1"
dependencies = [
"cast",
"itertools 0.10.5",
]
[[package]]
name = "critical-section"
version = "1.2.0"
@@ -2319,7 +2234,6 @@ dependencies = [
"aes-gcm",
"anyhow",
"arc-swap",
"ariadne",
"async-recursion",
"async-ringbuf",
"async-stream",
@@ -2341,7 +2255,6 @@ dependencies = [
"clap_complete",
"clap_complete_nushell",
"console-subscriber",
"criterion",
"crossbeam",
"ctor 0.8.0",
"dashmap",
@@ -2359,7 +2272,7 @@ dependencies = [
"gethostname 0.5.0",
"git-version",
"globwalk",
"guarden 0.2.0",
"guarden",
"hickory-client",
"hickory-proto",
"hickory-resolver",
@@ -2404,9 +2317,8 @@ dependencies = [
"prost-reflect",
"prost-reflect-build",
"prost-wkt-types",
"quanta",
"quinn",
"quinn-proto",
"quinn-plaintext",
"quote",
"rand 0.8.5",
"rcgen",
@@ -2418,7 +2330,6 @@ dependencies = [
"rstest",
"rust-i18n",
"rustls",
"seahash",
"serde",
"serde_json",
"serial_test",
@@ -2546,7 +2457,7 @@ dependencies = [
"dashmap",
"easytier",
"futures",
"guarden 0.1.2",
"guarden",
"jsonwebtoken",
"mimalloc",
"mockall",
@@ -3682,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",
]
@@ -3708,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"
@@ -4572,17 +4460,6 @@ dependencies = [
"once_cell",
]
[[package]]
name = "is-terminal"
version = "0.4.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46"
dependencies = [
"hermit-abi",
"libc",
"windows-sys 0.61.2",
]
[[package]]
name = "is-wsl"
version = "0.4.0"
@@ -5701,7 +5578,7 @@ version = "0.7.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "680998035259dcfcafe653688bf2aa6d3e2dc05e98be6ab46afb089dc84f1df8"
dependencies = [
"proc-macro-crate 3.5.0",
"proc-macro-crate 3.2.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -5943,12 +5820,6 @@ dependencies = [
"portable-atomic",
]
[[package]]
name = "oorandom"
version = "11.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e"
[[package]]
name = "opaque-debug"
version = "0.3.1"
@@ -6579,34 +6450,6 @@ dependencies = [
"time",
]
[[package]]
name = "plotters"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747"
dependencies = [
"num-traits",
"plotters-backend",
"plotters-svg",
"wasm-bindgen",
"web-sys",
]
[[package]]
name = "plotters-backend"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a"
[[package]]
name = "plotters-svg"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670"
dependencies = [
"plotters-backend",
]
[[package]]
name = "pnet"
version = "0.35.0"
@@ -6867,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]]
@@ -7019,12 +6862,9 @@ version = "0.16.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "590aa145fee8f7a26b5a6055365e7c5e89a5c1caae9869de76ec0ee73181a2f9"
dependencies = [
"base64 0.22.1",
"prost 0.14.3",
"prost-reflect-derive",
"prost-types 0.14.3",
"serde",
"serde-value",
]
[[package]]
@@ -7138,21 +6978,6 @@ version = "0.1.28"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b5a041e753da8b807c9255f28de81879c78c876392ff2469cde94799b2896b9d"
[[package]]
name = "quanta"
version = "0.12.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f3ab5a9d756f0d97bdc89019bd2e4ea098cf9cde50ee7564dde6b81ccc8f06c7"
dependencies = [
"crossbeam-utils",
"libc",
"once_cell",
"raw-cpuid",
"wasi 0.11.0+wasi-snapshot-preview1",
"web-sys",
"winapi",
]
[[package]]
name = "quick-error"
version = "2.0.1"
@@ -7197,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"
@@ -7406,15 +7243,6 @@ dependencies = [
"rand_core 0.5.1",
]
[[package]]
name = "raw-cpuid"
version = "11.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186"
dependencies = [
"bitflags 2.8.0",
]
[[package]]
name = "raw-window-handle"
version = "0.6.2"
@@ -7780,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",
@@ -9985,16 +9813,6 @@ dependencies = [
"zerovec",
]
[[package]]
name = "tinytemplate"
version = "1.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc"
dependencies = [
"serde",
"serde_json",
]
[[package]]
name = "tinyvec"
version = "1.8.0"
@@ -10114,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",
@@ -10189,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"
@@ -10235,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"
@@ -12094,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"
@@ -12478,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",
@@ -12715,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;
}
+13 -3
View File
@@ -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,
+2 -2
View File
@@ -5,8 +5,8 @@
"private": true,
"packageManager": "pnpm@9.12.1+sha512.e5a7e52a4183a02d5931057f7a0dbff9d5e9ce3161e33fa68ae392125b79282a8a8a470a51dfc8a0ed86221442eb2fb57019b0990ed24fab519bf0e1bc5ccfc4",
"scripts": {
"dev": "pnpm --dir ../easytier-web/frontend-lib build && vite",
"build": "pnpm --dir ../easytier-web/frontend-lib build && vue-tsc --noEmit && vite build",
"dev": "vite",
"build": "vue-tsc --noEmit && vite build",
"preview": "vite preview",
"tauri": "tauri",
"lint": "eslint . --ignore-pattern src-tauri",
+14 -11
View File
@@ -163,10 +163,14 @@ async function registerVpnServiceListener() {
)
}
function getRoutesForVpn(routes: Route[] | undefined, node_config: NetworkTypes.NetworkConfig): string[] {
function getRoutesForVpn(routes: Route[], node_config: NetworkTypes.NetworkConfig): string[] {
if (!routes) {
return []
}
const ret = []
for (const r of routes ?? []) {
for (let cidr of r.proxy_cidrs ?? []) {
for (const r of routes) {
for (let cidr of r.proxy_cidrs) {
if (!cidr.includes('/')) {
cidr += '/32'
}
@@ -174,9 +178,9 @@ function getRoutesForVpn(routes: Route[] | undefined, node_config: NetworkTypes.
}
}
for (const route of node_config.routes ?? []) {
ret.push(route)
}
node_config.routes.forEach(r => {
ret.push(r)
})
if (node_config.enable_magic_dns) {
ret.push('100.100.100.101/32')
@@ -211,15 +215,14 @@ export async function onNetworkInstanceChange(instanceId: string) {
console.log('vpn service skipped because no_tun is enabled', instanceId)
return
}
const curNetworkInfo = (await collectNetworkInfo(instanceId))?.info?.map?.[instanceId]
const curNetworkInfo = (await collectNetworkInfo(instanceId)).info.map[instanceId]
if (!curNetworkInfo || curNetworkInfo?.error_msg?.length) {
console.warn('vpn service skipped because network info is unavailable', instanceId, curNetworkInfo?.error_msg)
await doStopVpn()
return
}
const virtualIpv4 = curNetworkInfo.my_node_info?.virtual_ipv4
const virtual_ip = virtualIpv4?.address?.addr ? Utils.ipv4ToString(virtualIpv4.address) : undefined
const virtual_ip = Utils.ipv4ToString(curNetworkInfo?.my_node_info?.virtual_ipv4.address)
if (config.dhcp && (!virtual_ip || !virtual_ip.length)) {
console.log('DHCP enabled but no IP yet, will retry in', DHCP_POLLING_INTERVAL, 'ms')
@@ -234,7 +237,7 @@ export async function onNetworkInstanceChange(instanceId: string) {
return
}
let network_length = virtualIpv4?.network_length
let network_length = curNetworkInfo?.my_node_info?.virtual_ipv4.network_length
if (!network_length) {
network_length = 24
}
@@ -287,7 +290,7 @@ async function isNoTunEnabled(instanceId: string | undefined) {
async function findRunningTunInstanceId() {
const instanceIds = await listNetworkInstanceIds()
const runningIds = (instanceIds.running_inst_ids ?? []).map(Utils.UuidToStr)
const runningIds = instanceIds.running_inst_ids.map(Utils.UuidToStr)
console.log('vpn service sync running instances', JSON.stringify(runningIds))
for (const instanceId of runningIds) {
+2 -2
View File
@@ -9,7 +9,7 @@ export class GUIRemoteClient implements Api.RemoteClient {
await backend.runNetworkInstance(config, save);
}
async get_network_info(inst_id: string): Promise<NetworkTypes.NetworkInstanceRunningInfo | undefined> {
return backend.collectNetworkInfo(inst_id).then(infos => infos.info?.map?.[inst_id]);
return backend.collectNetworkInfo(inst_id).then(infos => infos.info.map[inst_id]);
}
async list_network_instance_ids(): Promise<Api.ListNetworkInstanceIdResponse> {
return backend.listNetworkInstanceIds();
@@ -44,4 +44,4 @@ export class GUIRemoteClient implements Api.RemoteClient {
return await backend.getNetworkMetas(instance_ids);
}
}
}
+5 -16
View File
@@ -13,18 +13,12 @@
"./*.css": "./dist/*.css"
},
"scripts": {
"codegen:proto": "node scripts/codegen-proto.mjs",
"dev": "pnpm codegen:proto && vite",
"build": "pnpm codegen:proto && vue-tsc -b && vite build",
"test": "pnpm test:config-ui && pnpm test:network-config",
"test:config-ui": "pnpm codegen:proto && vitest run --config vitest.config.ts",
"test:network-config": "pnpm build && node scripts/test-network-config.mjs",
"dev": "vite",
"build": "vue-tsc -b && vite build",
"preview": "vite preview"
},
"dependencies": {
"@primeuix/themes": "^1.2.3",
"@protobuf-ts/runtime": "2.11.1",
"@protobuf-ts/runtime-rpc": "2.11.1",
"@vueuse/core": "^11.1.0",
"axios": "^1.13.5",
"chart.js": "^4.5.0",
@@ -39,13 +33,9 @@
},
"devDependencies": {
"@modyfi/vite-plugin-yaml": "^1.1.0",
"@protobuf-ts/plugin": "2.11.1",
"@protobuf-ts/protoc": "2.11.1",
"@types/node": "^22.8.6",
"@vitejs/plugin-vue": "^5.1.4",
"@vue/test-utils": "^2.4.11",
"autoprefixer": "^10.4.20",
"happy-dom": "16.8.1",
"postcss": "^8.4.47",
"postcss-import": "^16.1.0",
"postcss-nested": "^7.0.2",
@@ -53,11 +43,10 @@
"typescript": "~5.6.3",
"vite": "^5.4.21",
"vite-plugin-dts": "^4.3.0",
"vitest": "^2.1.9",
"vue-tsc": "^2.1.10"
},
"peerDependencies": {
"primevue": "^4.3.9",
"vue": "^3.5.12"
"vue": "^3.5.12",
"primevue": "^4.3.9"
}
}
}
+10 -2482
View File
File diff suppressed because it is too large Load Diff
@@ -1,121 +0,0 @@
import { spawnSync } from 'node:child_process'
import { existsSync, mkdirSync, mkdtempSync, readdirSync, renameSync, rmSync, statSync } from 'node:fs'
import { createRequire } from 'node:module'
import { delimiter, dirname, resolve } from 'node:path'
import { fileURLToPath } from 'node:url'
const require = createRequire(import.meta.url)
const root = resolve(dirname(fileURLToPath(import.meta.url)), '..')
const protoRoot = resolve(root, '../../easytier/src/proto')
const generatedRoot = resolve(root, 'src/generated')
const outDir = resolve(generatedRoot, 'proto')
const nodeBinDir = resolve(root, 'node_modules/.bin')
const protocWrapper = require.resolve('@protobuf-ts/protoc/protoc.js')
const protobufTsPluginRoot = dirname(require.resolve('@protobuf-ts/plugin/package.json'))
const protoFiles = [
'common.proto',
'acl.proto',
'api_instance.proto',
'api_manage.proto',
'peer_rpc.proto',
'error.proto',
]
function installGeneratedFiles(fromDir, toDir) {
mkdirSync(toDir, { recursive: true })
for (const entry of readdirSync(fromDir)) {
const source = resolve(fromDir, entry)
const target = resolve(toDir, entry)
if (statSync(source).isDirectory()) {
installGeneratedFiles(source, target)
continue
}
renameSync(source, target)
}
}
function findExecutableInPath(command, extensions = ['']) {
const envPath = process.env[pathEnvKey()]
if (typeof envPath !== 'string') return undefined
const nodeBinSuffix = ['node_modules/.bin', 'node_modules\\.bin']
for (const entry of envPath.split(delimiter)) {
if (!entry || nodeBinSuffix.some((suffix) => entry.endsWith(suffix))) continue
for (const extension of extensions) {
const candidate = resolve(entry, `${command}${extension}`)
if (existsSync(candidate)) return candidate
}
}
return undefined
}
function pathEnvKey() {
return Object.keys(process.env).find((key) => key.toLowerCase() === 'path') ?? 'PATH'
}
function withNodeBinPath() {
const key = pathEnvKey()
const currentPath = process.env[key]
return {
...process.env,
[key]: currentPath ? `${nodeBinDir}${delimiter}${currentPath}` : nodeBinDir,
}
}
function getProtocCommand() {
const extensions = process.platform === 'win32' ? ['.exe'] : ['']
const systemProtoc = findExecutableInPath('protoc', extensions)
if (systemProtoc) {
return {
command: systemProtoc,
argsPrefix: ['--proto_path', protobufTsPluginRoot],
}
}
return {
command: process.execPath,
argsPrefix: [protocWrapper],
}
}
mkdirSync(generatedRoot, { recursive: true })
const tmpDir = mkdtempSync(resolve(generatedRoot, '.proto-'))
const protocCommand = getProtocCommand()
try {
const result = spawnSync(protocCommand.command, [
...protocCommand.argsPrefix,
'-I',
protoRoot,
`--ts_out=${tmpDir}`,
'--ts_opt=use_proto_field_name,server_none,client_none,ts_nocheck',
...protoFiles.map((file) => resolve(protoRoot, file)),
], {
cwd: root,
env: withNodeBinPath(),
stdio: 'inherit',
shell: false,
})
if (result.error) {
throw result.error
}
const status = result.status ?? 1
if (status === 0) {
installGeneratedFiles(tmpDir, outDir)
}
process.exit(status)
} finally {
rmSync(tmpDir, { recursive: true, force: true })
}
@@ -1,540 +0,0 @@
import assert from 'node:assert/strict'
import fs from 'node:fs'
import path from 'node:path'
import { fileURLToPath, pathToFileURL } from 'node:url'
import ts from 'typescript'
const __dirname = path.dirname(fileURLToPath(import.meta.url))
const projectRoot = path.resolve(__dirname, '..')
const generatedApiManagePath = path.join(projectRoot, 'src/generated/proto/api_manage.ts')
const distPath = path.join(projectRoot, 'dist/easytier-frontend-lib.js')
const { NetworkTypes } = await import(pathToFileURL(distPath))
const {
AclAction,
AclChainType,
AclProtocol,
CompressionAlgoPb,
DEFAULT_NETWORK_CONFIG,
NetworkingMethod,
normalizeNetworkConfig,
toBackendNetworkConfig,
} = NetworkTypes
const BOOLEAN_CONFIG_FIELDS = [
'dhcp',
'enable_vpn_portal',
'advanced_settings',
'latency_first',
'use_smoltcp',
'disable_ipv6',
'enable_kcp_proxy',
'disable_kcp_input',
'disable_p2p',
'bind_device',
'no_tun',
'enable_exit_node',
'relay_all_peer_rpc',
'multi_thread',
'enable_relay_network_whitelist',
'enable_manual_routes',
'proxy_forward_by_system',
'disable_encryption',
'enable_socks5',
'disable_udp_hole_punching',
'enable_magic_dns',
'enable_private_mode',
'enable_quic_proxy',
'disable_quic_input',
'disable_sym_hole_punching',
'p2p_only',
'lazy_p2p',
'need_p2p',
'disable_upnp',
'ipv6_public_addr_provider',
'ipv6_public_addr_auto',
'disable_relay_data',
'enable_udp_broadcast_relay',
'disable_tcp_hole_punching',
]
function readGeneratedNetworkConfigFields() {
const source = ts.createSourceFile(
generatedApiManagePath,
fs.readFileSync(generatedApiManagePath, 'utf8'),
ts.ScriptTarget.Latest,
true,
)
for (const statement of source.statements) {
if (!ts.isInterfaceDeclaration(statement) || statement.name.text !== 'NetworkConfig') {
continue
}
return statement.members
.filter(ts.isPropertySignature)
.map((member) => member.name.getText(source).replace(/^['"]|['"]$/g, ''))
}
throw new Error(`NetworkConfig interface not found in ${generatedApiManagePath}`)
}
function expectNoCamelCaseKeys(value, pathSegments = []) {
if (!value || typeof value !== 'object') {
return
}
if (Array.isArray(value)) {
value.forEach((item, index) => expectNoCamelCaseKeys(item, [...pathSegments, String(index)]))
return
}
for (const [key, child] of Object.entries(value)) {
assert.equal(
/[A-Z]/.test(key),
false,
`JSON key should use proto field name: ${[...pathSegments, key].join('.')}`,
)
expectNoCamelCaseKeys(child, [...pathSegments, key])
}
}
function allFieldFixture() {
return {
...DEFAULT_NETWORK_CONFIG(),
instance_id: '11111111-2222-3333-4444-555555555555',
dhcp: false,
virtual_ipv4: '10.9.8.7',
network_length: 25,
hostname: 'frontend-e2e',
network_name: 'full-field-network',
network_secret: 'full-field-secret',
networking_method: NetworkingMethod.Manual,
public_server_url: 'tcp://public.example:11010',
peer_urls: [' tcp://peer-a:11010 ', '', 'udp://peer-b:11010'],
peers: [
{
uri: 'tcp://peer-a:11010',
peer_public_key: 'peer-a-public-key',
},
],
proxy_cidrs: ['10.10.0.0/16', '192.168.2.0/24->10.99.0.0/24'],
enable_vpn_portal: true,
vpn_portal_listen_port: 23000,
vpn_portal_client_network_addr: '10.88.0.0',
vpn_portal_client_network_len: 24,
advanced_settings: true,
listener_urls: ['tcp://0.0.0.0:12010', 'udp://0.0.0.0:12010'],
latency_first: true,
dev_name: 'et-full',
use_smoltcp: true,
disable_ipv6: true,
enable_kcp_proxy: true,
disable_kcp_input: true,
disable_p2p: true,
bind_device: false,
no_tun: true,
enable_exit_node: true,
relay_all_peer_rpc: true,
multi_thread: false,
enable_relay_network_whitelist: true,
relay_network_whitelist: ['10.0.0.0/8', 'fd00::/8'],
enable_manual_routes: true,
routes: ['10.20.0.0/16', 'fd00:20::/64'],
exit_nodes: ['10.9.8.1', 'fd00::1'],
proxy_forward_by_system: true,
disable_encryption: true,
enable_socks5: true,
socks5_port: 1081,
disable_udp_hole_punching: true,
mtu: 1280,
mapped_listeners: ['tcp://127.0.0.1:13010'],
enable_magic_dns: true,
enable_private_mode: true,
enable_quic_proxy: true,
disable_quic_input: true,
quic_listen_port: 14010,
port_forwards: [
{
proto: 'tcp',
bind_ip: '127.0.0.1',
bind_port: 8080,
dst_ip: '10.9.8.7',
dst_port: 80,
},
{
proto: 'udp',
bind_ip: '0.0.0.0',
bind_port: 5353,
dst_ip: '10.9.8.8',
dst_port: 53,
},
],
disable_sym_hole_punching: true,
p2p_only: true,
data_compress_algo: CompressionAlgoPb.Zstd,
encryption_algorithm: 'aes-gcm',
disable_tcp_hole_punching: true,
secure_mode: {
enabled: true,
local_private_key: 'private-key',
local_public_key: 'public-key',
},
acl: {
acl_v1: {
group: {
declares: [
{
group_name: 'ops',
group_secret: 'ops-secret',
},
],
members: ['node-a', 'node-b'],
},
chains: [
{
name: 'forward-chain',
chain_type: AclChainType.Forward,
description: 'forward traffic',
enabled: true,
default_action: AclAction.Drop,
rules: [
{
name: 'allow-web',
description: 'allow web traffic',
priority: 100,
enabled: true,
protocol: AclProtocol.TCP,
ports: ['80', '443'],
source_ips: ['10.0.0.0/8'],
destination_ips: ['10.9.8.7/32'],
source_ports: ['1024-65535'],
action: AclAction.Allow,
rate_limit: 1000,
burst_limit: 2000,
stateful: true,
source_groups: ['ops'],
destination_groups: ['web'],
},
],
},
],
},
},
credential_file: '/tmp/easytier-credential.toml',
lazy_p2p: true,
need_p2p: true,
instance_recv_bps_limit: '9007199254740993',
disable_upnp: true,
ipv6_public_addr_provider: true,
ipv6_public_addr_auto: true,
ipv6_public_addr_prefix: '2001:db8:1::/64',
disable_relay_data: true,
enable_udp_broadcast_relay: true,
socket_mark: 1234,
}
}
function assertFixtureCoversGeneratedFields() {
const generatedFields = readGeneratedNetworkConfigFields()
const fixtureFields = new Set(Object.keys(allFieldFixture()))
const missing = generatedFields.filter((field) => !fixtureFields.has(field))
assert.deepEqual(missing, [], 'all generated NetworkConfig fields should be represented in the fixture')
}
function assertFullFieldRoundTrip() {
const input = allFieldFixture()
const normalized = normalizeNetworkConfig(input)
assert.equal(normalized.peer_urls.join(','), 'tcp://peer-a:11010,udp://peer-b:11010')
assert.equal(normalized.instance_recv_bps_limit, '9007199254740993')
assert.equal(normalized.data_compress_algo, CompressionAlgoPb.Zstd)
assert.equal(normalized.acl.acl_v1.chains[0].chain_type, AclChainType.Forward)
assert.equal(normalized.acl.acl_v1.chains[0].rules[0].protocol, AclProtocol.TCP)
const backend = toBackendNetworkConfig(normalized)
expectNoCamelCaseKeys(backend)
for (const field of readGeneratedNetworkConfigFields()) {
assert.ok(field in backend, `backend JSON should include fixture field ${field}`)
}
assert.equal(backend.networking_method, 'Manual')
assert.equal(backend.public_server_url, '')
assert.deepEqual(backend.peer_urls, ['tcp://peer-a:11010', 'udp://peer-b:11010'])
assert.equal(backend.peers[0].peer_public_key, 'peer-a-public-key')
assert.deepEqual(backend.peers[1], { uri: 'udp://peer-b:11010' })
assert.equal(backend.data_compress_algo, 'Zstd')
assert.equal(backend.instance_recv_bps_limit, '9007199254740993')
assert.equal(backend.secure_mode.enabled, true)
assert.equal(backend.secure_mode.local_private_key, 'private-key')
assert.equal(backend.acl.acl_v1.chains[0].chain_type, 'Forward')
assert.equal(backend.acl.acl_v1.chains[0].default_action, 'Drop')
assert.equal(backend.acl.acl_v1.chains[0].rules[0].protocol, 'TCP')
assert.equal(backend.acl.acl_v1.chains[0].rules[0].action, 'Allow')
assert.equal(backend.port_forwards[1].proto, 'udp')
assert.equal(backend.socket_mark, 1234)
}
function assertBooleanFieldValuesPreserved() {
const input = allFieldFixture()
const normalized = normalizeNetworkConfig(input)
const backend = toBackendNetworkConfig(normalized)
for (const field of BOOLEAN_CONFIG_FIELDS) {
assert.equal(
normalized[field],
input[field],
`normalized config should preserve boolean field ${field}`,
)
assert.equal(
backend[field],
input[field],
`backend JSON should preserve boolean field ${field}`,
)
}
}
function assertEnumCompatibility() {
const normalized = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
networking_method: 'Manual',
data_compress_algo: 'Zstd',
acl: {
acl_v1: {
group: { declares: [], members: [] },
chains: [
{
chain_type: 'Forward',
default_action: 'Drop',
rules: [
{
protocol: 'TCP',
action: 'Allow',
},
],
},
],
},
},
})
assert.equal(normalized.data_compress_algo, CompressionAlgoPb.Zstd)
assert.equal(normalized.acl.acl_v1.chains[0].chain_type, AclChainType.Forward)
assert.equal(normalized.acl.acl_v1.chains[0].default_action, AclAction.Drop)
assert.equal(normalized.acl.acl_v1.chains[0].rules[0].protocol, AclProtocol.TCP)
assert.equal(normalized.acl.acl_v1.chains[0].rules[0].action, AclAction.Allow)
const backend = toBackendNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
data_compress_algo: 'Zstd',
acl: {
acl_v1: {
group: { declares: [], members: [] },
chains: [
{
chain_type: 'Forward',
default_action: 'Drop',
rules: [
{
protocol: 'TCP',
action: 'Allow',
},
],
},
],
},
},
})
assert.equal(backend.data_compress_algo, 'Zstd')
assert.equal(backend.acl.acl_v1.chains[0].chain_type, 'Forward')
assert.equal(backend.acl.acl_v1.chains[0].rules[0].protocol, 'TCP')
}
function assertAclDefaultsAndExplicitZero() {
const partialAcl = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
acl: {
acl_v1: {
group: { declares: [], members: [] },
chains: [{ rules: [{}] }],
},
},
})
const defaultedChain = partialAcl.acl.acl_v1.chains[0]
assert.equal(defaultedChain.chain_type, AclChainType.UnspecifiedChain)
assert.equal(defaultedChain.default_action, AclAction.Allow)
assert.equal(defaultedChain.rules[0].protocol, AclProtocol.Any)
assert.equal(defaultedChain.rules[0].action, AclAction.Allow)
const explicitZero = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
acl: {
acl_v1: {
group: { declares: [], members: [] },
chains: [
{
chain_type: 0,
default_action: 0,
rules: [{ protocol: 0, action: 0 }],
},
],
},
},
})
const zeroChain = explicitZero.acl.acl_v1.chains[0]
assert.equal(zeroChain.chain_type, AclChainType.UnspecifiedChain)
assert.equal(zeroChain.default_action, AclAction.Noop)
assert.equal(zeroChain.rules[0].protocol, AclProtocol.Unspecified)
assert.equal(zeroChain.rules[0].action, AclAction.Noop)
}
function assertNetworkingMethodNormalization() {
const publicServer = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
networking_method: 'PublicServer',
public_server_url: ' tcp://public.example:11010 ',
peer_urls: ['tcp://manual.example:11010'],
})
assert.equal(publicServer.networking_method, NetworkingMethod.Manual)
assert.equal(publicServer.public_server_url, '')
assert.deepEqual(publicServer.peer_urls, ['tcp://public.example:11010'])
const standalone = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
networking_method: 'Standalone',
peer_urls: ['tcp://manual.example:11010'],
})
assert.equal(standalone.networking_method, NetworkingMethod.Manual)
assert.deepEqual(standalone.peer_urls, [])
const missing = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
networking_method: undefined,
peer_urls: [' tcp://one ', '', 'udp://two '],
})
assert.deepEqual(missing.peer_urls, ['tcp://one', 'udp://two'])
const publicServerMissingUrl = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
networking_method: 'PublicServer',
public_server_url: '',
peer_urls: ['tcp://manual.example:11010'],
})
assert.deepEqual(publicServerMissingUrl.peer_urls, [])
}
function assertPeerPublicKeysPreserved() {
const normalized = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
peer_urls: [],
peers: [
{
uri: ' tcp://peer-a:11010 ',
peer_public_key: 'peer-a-public-key',
},
],
})
assert.deepEqual(normalized.peer_urls, ['tcp://peer-a:11010'])
assert.deepEqual(normalized.peers, [
{
uri: 'tcp://peer-a:11010',
peer_public_key: 'peer-a-public-key',
},
])
const unchangedUrl = toBackendNetworkConfig({
...normalized,
peer_urls: ['tcp://peer-a:11010', 'tcp://peer-b:11010'],
})
assert.equal(unchangedUrl.peers[0].peer_public_key, 'peer-a-public-key')
assert.deepEqual(unchangedUrl.peers[1], { uri: 'tcp://peer-b:11010' })
const changedUrl = toBackendNetworkConfig({
...normalized,
peer_urls: ['tcp://peer-c:11010'],
})
assert.deepEqual(changedUrl.peers, [{ uri: 'tcp://peer-c:11010' }])
const clearedUrls = toBackendNetworkConfig({
...normalized,
peer_urls: [],
})
assert.deepEqual(clearedUrls.peer_urls ?? [], [])
assert.deepEqual(clearedUrls.peers ?? [], [])
}
function assertNumberBoundaries() {
const safeLimit = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: '12345',
})
assert.equal(safeLimit.instance_recv_bps_limit, 12345)
const largeLimit = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: '9007199254740993',
})
assert.equal(largeLimit.instance_recv_bps_limit, '9007199254740993')
assert.equal(toBackendNetworkConfig(largeLimit).instance_recv_bps_limit, '9007199254740993')
const invalidNumbers = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
mtu: Number.NaN,
instance_recv_bps_limit: Number.POSITIVE_INFINITY,
})
assert.equal(invalidNumbers.mtu, null)
assert.equal(invalidNumbers.instance_recv_bps_limit, null)
const emptyLimit = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: '',
})
assert.equal(emptyLimit.instance_recv_bps_limit, null)
const zeroLimit = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: '0',
})
assert.equal(zeroLimit.instance_recv_bps_limit, null)
assert.equal(toBackendNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: 0,
}).instance_recv_bps_limit, undefined)
const oversizedLimit = normalizeNetworkConfig({
...DEFAULT_NETWORK_CONFIG(),
instance_recv_bps_limit: '18446744073709551616',
})
assert.equal(oversizedLimit.instance_recv_bps_limit, null)
}
const tests = [
assertFixtureCoversGeneratedFields,
assertFullFieldRoundTrip,
assertBooleanFieldValuesPreserved,
assertEnumCompatibility,
assertAclDefaultsAndExplicitZero,
assertNetworkingMethodNormalization,
assertPeerPublicKeysPreserved,
assertNumberBoundaries,
]
for (const test of tests) {
test()
console.log(`ok ${test.name}`)
}
@@ -9,7 +9,7 @@ import {
normalizeNetworkConfig,
removeRow
} from '../types/network'
import { computed, ref, onMounted, onUnmounted, watch } from 'vue'
import { ref, onMounted, onUnmounted, watch } from 'vue'
import { useI18n } from 'vue-i18n'
import AclManager from './acl/AclManager.vue'
import UrlListInput from './UrlListInput.vue'
@@ -134,7 +134,6 @@ function savePortForward() {
const portForwardContainer = ref<HTMLElement | null>(null);
const isCompact = ref(false);
const UINT64_MAX = (1n << 64n) - 1n
onMounted(() => {
if (portForwardContainer.value) {
@@ -162,39 +161,6 @@ function syncNormalizedNetwork(network: NetworkConfig | undefined): void {
}
watch(() => curNetwork.value, syncNormalizedNetwork, { immediate: true, deep: false })
function parseInstanceRecvBpsLimitInput(value: string): number | string | null | undefined {
const trimmed = value.trim()
if (trimmed.length === 0) {
return null
}
if (!/^\d+$/.test(trimmed)) {
return undefined
}
const limit = BigInt(trimmed)
if (limit === 0n) {
return null
}
if (limit > UINT64_MAX) {
return undefined
}
return limit <= BigInt(Number.MAX_SAFE_INTEGER) ? Number(limit) : limit.toString()
}
const instanceRecvBpsLimitInput = computed<string>({
get: () => {
const limit = curNetwork.value.instance_recv_bps_limit
return limit == null ? '' : String(limit)
},
set: (value) => {
const limit = parseInstanceRecvBpsLimitInput(value)
if (limit !== undefined) {
curNetwork.value.instance_recv_bps_limit = limit
}
},
})
</script>
<template>
@@ -351,9 +317,9 @@ const instanceRecvBpsLimitInput = computed<string>({
<span class="pi pi-question-circle ml-2 self-center"
v-tooltip="t('instance_recv_bps_limit_help')"></span>
</div>
<InputText id="instance_recv_bps_limit" v-model="instanceRecvBpsLimitInput"
aria-describedby="instance_recv_bps_limit-help" inputmode="numeric" pattern="[0-9]*"
:placeholder="t('instance_recv_bps_limit_placeholder')" fluid />
<InputNumber id="instance_recv_bps_limit" v-model="curNetwork.instance_recv_bps_limit"
aria-describedby="instance_recv_bps_limit-help" :format="false"
:placeholder="t('instance_recv_bps_limit_placeholder')" :min="1" fluid />
</div>
</div>
@@ -35,7 +35,7 @@ const currentNetworkConfig = ref<NetworkTypes.NetworkConfig | undefined>(undefin
const listInstanceIdResponse = ref<Api.ListNetworkInstanceIdResponse | undefined>(undefined);
const isRunning = (instanceId: string) => {
return (listInstanceIdResponse.value?.running_inst_ids ?? []).map(Utils.UuidToStr).includes(instanceId);
return listInstanceIdResponse.value?.running_inst_ids.map(Utils.UuidToStr).includes(instanceId);
}
const networkMetaCache = ref<Record<string, Api.NetworkMeta>>({});
@@ -46,7 +46,7 @@ const loadNetworkMetas = async (instanceIds: string[]) => {
try {
const response = await props.api.get_network_metas(missingIds);
Object.assign(networkMetaCache.value, response.metas ?? {});
Object.assign(networkMetaCache.value, response.metas);
} catch (e) {
console.error("Failed to load network metas", e);
}
@@ -80,8 +80,8 @@ const updateInstanceList = () => {
let insts = new Set<string>();
let t = listInstanceIdResponse.value;
if (t) {
(t.running_inst_ids ?? []).forEach((u) => insts.add(Utils.UuidToStr(u)));
(t.disabled_inst_ids ?? []).forEach((u) => insts.add(Utils.UuidToStr(u)));
t.running_inst_ids.forEach((u) => insts.add(Utils.UuidToStr(u)));
t.disabled_inst_ids.forEach((u) => insts.add(Utils.UuidToStr(u)));
}
const newList = Array.from(insts).map((instance: string) => {
@@ -149,7 +149,7 @@ const networkIsDisabled = computed(() => {
if (!selectedInstanceId.value) {
return false;
}
return (listInstanceIdResponse.value?.disabled_inst_ids ?? []).map(Utils.UuidToStr).includes(selectedInstanceId.value?.uuid);
return listInstanceIdResponse.value?.disabled_inst_ids.map(Utils.UuidToStr).includes(selectedInstanceId.value?.uuid);
});
watch(networkIsDisabled, async (newVal, oldVal) => {
if (newVal !== oldVal && newVal === true) {
@@ -287,35 +287,17 @@ const loadNetworkInstanceIds = async () => {
}
const loadCurrentNetworkInfo = async () => {
const selected = selectedInstanceId.value?.uuid;
if (!selected) {
curNetworkInfo.value = null;
if (!selectedInstanceId.value) {
return;
}
if (!needShowNetworkStatus.value) {
curNetworkInfo.value = null;
return;
}
if (curNetworkInfo.value?.instance_id !== selected) {
curNetworkInfo.value = null;
}
let network_info = await props.api.get_network_info(selected);
if (selectedInstanceId.value?.uuid !== selected) {
return;
}
if (!network_info) {
curNetworkInfo.value = {
instance_id: selected,
running: false,
error_msg: t('web.device_management.network_info_unavailable'),
} as NetworkTypes.NetworkInstance;
return;
}
let network_info = await props.api.get_network_info(selectedInstanceId.value.uuid);
curNetworkInfo.value = {
instance_id: selected,
instance_id: selectedInstanceId.value.uuid,
running: network_info?.running ?? false,
error_msg: network_info?.error_msg ?? '',
detail: network_info,
@@ -510,7 +492,7 @@ onUnmounted(() => {
<div class="flex items-center min-w-0">
<div class="mr-4 min-w-0 flex-1">
<span class="truncate block">{{ t('network_name') }}: {{
slotProps.option.meta?.network_name ?? slotProps.option.uuid }}</span>
slotProps.option.meta.network_name }}</span>
</div>
<Tag class="my-auto leading-3 shrink-0"
:severity="isRunning(slotProps.option.uuid) ? 'success' : 'info'"
@@ -587,13 +569,10 @@ onUnmounted(() => {
<h2 class="text-xl font-medium">{{ t('web.device_management.network_status') }}</h2>
</div>
<Status v-if="curNetworkInfo && curNetworkInfo.error_msg === ''" v-bind:cur-network-inst="curNetworkInfo"
<Status v-if="(curNetworkInfo?.error_msg ?? '') === ''" v-bind:cur-network-inst="curNetworkInfo"
class="mb-4">
</Status>
<Message v-else-if="curNetworkInfo?.error_msg" severity="error" class="mb-4">{{
curNetworkInfo.error_msg }}</Message>
<Message v-else severity="info" class="mb-4">{{ t('web.device_management.loading_network_status') }}
</Message>
<Message v-else severity="error" class="mb-4">{{ curNetworkInfo?.error_msg }}</Message>
<div class="text-center mt-4">
<Button @click="stopNetwork" :disabled="!currentNetworkControl.deletable.value"
@@ -1,10 +1,10 @@
<script setup lang="ts">
import { useTimeAgo } from '@vueuse/core'
import { IPv4 } from 'ip-num/IPNumber'
import { NetworkInstance, type TunnelInfo, type NodeInfo, type PeerRoutePair } from '../types/network'
import { useI18n } from 'vue-i18n';
import { computed, onMounted, onUnmounted, ref } from 'vue';
import { ipv4InetToString, ipv4ToString, ipv6ToString } from '../modules/utils';
import { latencyMs, lossRate, numericValue, peerConns } from '../modules/statusDisplay';
import { Badge, DataTable, Column, Tag, Chip, Button, Dialog, ScrollPanel, Timeline, Divider, Card, } from 'primevue';
import NetworkChart from './NetworkChart.vue';
@@ -39,8 +39,8 @@ function routeCost(info: any) {
return '?'
}
function resolveObjPath(path: string, obj: any = globalThis, separator = '.') {
const properties = path.split(separator)
function resolveObjPath(path: string, obj = globalThis, separator = '.') {
const properties = Array.isArray(path) ? path : path.split(separator)
return properties.reduce((prev, curr) => prev?.[curr], obj)
}
@@ -48,17 +48,10 @@ function statsCommon(info: any, field: string): number | undefined {
if (!info.peer)
return undefined
let sum = 0
let hasValue = false
for (const conn of peerConns(info)) {
const value = numericValue(resolveObjPath(field, conn))
if (value === undefined)
continue
sum += value
hasValue = true
}
return hasValue ? sum : undefined
const conns = info.peer.conns
return conns.reduce((acc: number, conn: any) => {
return acc + resolveObjPath(field, conn)
}, 0)
}
function humanFileSize(bytes: number, si = false, dp = 1) {
@@ -81,6 +74,14 @@ function humanFileSize(bytes: number, si = false, dp = 1) {
return `${bytes.toFixed(dp)} ${units[u]}`
}
function latencyMs(info: PeerRoutePair) {
let lat_us_sum = statsCommon(info, 'stats.latency_us')
if (lat_us_sum === undefined)
return ''
lat_us_sum = lat_us_sum / 1000 / info.peer!.conns.length
return `${lat_us_sum % 1 > 0 ? Math.round(lat_us_sum) + 1 : Math.round(lat_us_sum)}ms`
}
function txBytes(info: PeerRoutePair) {
const tx = statsCommon(info, 'stats.tx_bytes')
return tx ? humanFileSize(tx) : ''
@@ -91,6 +92,11 @@ function rxBytes(info: PeerRoutePair) {
return rx ? humanFileSize(rx) : ''
}
function lossRate(info: PeerRoutePair) {
const lossRate = statsCommon(info, 'loss_rate')
return lossRate !== undefined ? `${Math.round(lossRate * 100)}%` : ''
}
function version(info: PeerRoutePair) {
return info.route.version === '' ? 'unknown' : info.route.version
}
@@ -99,7 +105,7 @@ function ipFormat(info: PeerRoutePair) {
const ip = info.route.ipv4_addr
if (typeof ip === 'string')
return ip
return ip ? ipv4InetToString(ip) : ''
return ip ? `${IPv4.fromNumber(ip.address.addr)}/${ip.network_length}` : ''
}
function oneTunnelProto(tunnel?: TunnelInfo): string {
@@ -125,7 +131,7 @@ function oneTunnelProto(tunnel?: TunnelInfo): string {
}
function tunnelProto(info: PeerRoutePair) {
return [...new Set(peerConns(info).map(c => oneTunnelProto(c.tunnel)))].join(',')
return [...new Set(info.peer?.conns.map(c => oneTunnelProto(c.tunnel)))].join(',')
}
const myNodeInfo = computed(() => {
@@ -200,7 +206,7 @@ const myNodeInfoChips = computed(() => {
// local ipv4s
const local_ipv4s = my_node_info.ips?.interface_ipv4s
for (const [idx, ip] of local_ipv4s?.entries() ?? []) {
for (const [idx, ip] of local_ipv4s?.entries()) {
chips.push({
label: `Local IPv4 ${idx}: ${ipv4ToString(ip)}`,
icon: '',
@@ -209,7 +215,7 @@ const myNodeInfoChips = computed(() => {
// local ipv6s
const local_ipv6s = my_node_info.ips?.interface_ipv6s
for (const [idx, ip] of local_ipv6s?.entries() ?? []) {
for (const [idx, ip] of local_ipv6s?.entries()) {
chips.push({
label: `Local IPv6 ${idx}: ${ipv6ToString(ip)}`,
icon: '',
@@ -220,7 +226,7 @@ const myNodeInfoChips = computed(() => {
const public_ip = my_node_info.ips?.public_ipv4
if (public_ip) {
chips.push({
label: `Public IP: ${ipv4ToString(public_ip)}`,
label: `Public IP: ${IPv4.fromNumber(public_ip.addr)}`,
icon: '',
} as Chip)
}
@@ -235,7 +241,7 @@ const myNodeInfoChips = computed(() => {
// listeners:
const listeners = my_node_info.listeners
for (const [idx, listener] of listeners?.entries() ?? []) {
for (const [idx, listener] of listeners?.entries()) {
chips.push({
label: `Listener ${idx}: ${listener.url}`,
icon: '',
@@ -282,14 +288,6 @@ function natType(info: PeerRoutePair): string {
return ''
}
function isPublicServerRoute(info: PeerRoutePair): boolean {
return info.route?.feature_flag?.is_public_server ?? false
}
function shouldAvoidRelayData(info: PeerRoutePair): boolean {
return info.route?.feature_flag?.avoid_relay_data ?? false
}
const peerCount = computed(() => {
if (!peerRouteInfos.value)
return 0
@@ -344,7 +342,7 @@ function showEventLogs() {
if (!detail)
return
dialogContent.value = detail.events?.map((event: string) => JSON.parse(event)) ?? []
dialogContent.value = detail.events.map((event: string) => JSON.parse(event))
dialogHeader.value = 'event_log'
dialogVisible.value = true
}
@@ -436,16 +434,16 @@ function showEventLogs() {
<Column :field="ipFormat" :header="t('virtual_ipv4')" />
<Column :header="t('hostname')">
<template #body="slotProps">
<div v-if="!slotProps.data.route.cost || !isPublicServerRoute(slotProps.data)"
<div v-if="!slotProps.data.route.cost || !slotProps.data.route.feature_flag.is_public_server"
v-tooltip="slotProps.data.route.hostname">
{{
slotProps.data.route.hostname }}
</div>
<div v-else v-tooltip="slotProps.data.route.hostname" class="space-x-1">
<Tag v-if="isPublicServerRoute(slotProps.data)" severity="info" value="Info">
<Tag v-if="slotProps.data.route.feature_flag.is_public_server" severity="info" value="Info">
{{ t('status.server') }}
</Tag>
<Tag v-if="shouldAvoidRelayData(slotProps.data)" severity="warn" value="Warn">
<Tag v-if="slotProps.data.route.feature_flag.avoid_relay_data" severity="warn" value="Warn">
{{ t('status.relay') }}
</Tag>
</div>
@@ -11,17 +11,8 @@ const props = defineProps<{
const list = defineModel<string[]>({ required: true })
const fallbackUrl = () => {
const protoKeys = Object.keys(props.protos)
const defaultProto = protoKeys.includes('tcp')
? 'tcp'
: (protoKeys[0] ?? 'tcp')
const defaultPort = props.protos[defaultProto] ?? 11010
return `${defaultProto}://0.0.0.0:${defaultPort}`
}
const addUrl = () => {
list.value.push(props.defaultUrl || fallbackUrl())
list.value.push(props.defaultUrl || 'tcp://0.0.0.0:11010')
}
const removeUrl = (index: number) => {
@@ -2,7 +2,7 @@
import { Button, Column, DataTable, Divider, InputText, Select, SelectButton, ToggleButton } from 'primevue'
import { ref, watch } from 'vue'
import { useI18n } from 'vue-i18n'
import { AclAction, AclChain, AclChainType, AclProtocol, AclRule, ensureAclChain, ensureAclRuleLists } from '../../types/network'
import { AclAction, AclChain, AclChainType, AclProtocol, AclRule } from '../../types/network'
import AclRuleDialog from './AclRuleDialog.vue'
const props = defineProps<{
@@ -13,11 +13,7 @@ const chain = defineModel<AclChain>({ required: true })
const { t } = useI18n()
function rules() {
return ensureAclChain(chain.value).rules
}
watch(() => rules(), (newRules) => {
watch(() => chain.value.rules, (newRules) => {
if (!newRules) return
const isSorted = newRules.every((rule, i) => i === 0 || (rule.priority || 0) <= (newRules[i - 1].priority || 0))
if (!isSorted) {
@@ -64,7 +60,7 @@ function addRule() {
editingRule.value = {
name: '',
description: '',
priority: rules().length,
priority: chain.value.rules.length,
enabled: true,
protocol: AclProtocol.Any,
ports: [],
@@ -83,31 +79,28 @@ function addRule() {
function editRule(index: number) {
editingRuleIndex.value = index
editingRule.value = ensureAclRuleLists(JSON.parse(JSON.stringify(rules()[index])))
editingRule.value = JSON.parse(JSON.stringify(chain.value.rules[index]))
showRuleDialog.value = true
}
function deleteRule(index: number) {
rules().splice(index, 1)
chain.value.rules.splice(index, 1)
}
function saveRule(rule: AclRule) {
const chainRules = rules()
ensureAclRuleLists(rule)
if (editingRuleIndex.value === -1) {
chainRules.push(rule)
chain.value.rules.push(rule)
} else {
chainRules[editingRuleIndex.value] = rule
chain.value.rules[editingRuleIndex.value] = rule
}
chainRules.sort((a, b) => (b.priority || 0) - (a.priority || 0))
chain.value.rules.sort((a, b) => (b.priority || 0) - (a.priority || 0))
}
function onRowReorder(event: any) {
chain.value.rules = event.value ?? []
const chainRules = rules()
chain.value.rules = event.value
// Update priorities based on new order (higher priority at top)
chainRules.forEach((rule, index) => {
rule.priority = chainRules.length - index - 1
chain.value.rules.forEach((rule, index) => {
rule.priority = chain.value.rules.length - index - 1
})
}
</script>
@@ -150,7 +143,7 @@ function onRowReorder(event: any) {
<Button icon="pi pi-plus" :label="t('acl.add_rule')" severity="success" size="small" @click="addRule" />
</div>
<DataTable :value="rules()" @row-reorder="onRowReorder" responsiveLayout="scroll">
<DataTable :value="chain.rules" @row-reorder="onRowReorder" responsiveLayout="scroll">
<Column rowReorder headerStyle="width: 3rem" />
<Column field="enabled" :header="t('acl.rule.enabled')">
<template #body="{ data }">
@@ -1,8 +1,8 @@
<script setup lang="ts">
import { Button, Column, DataTable, Dialog, InputText, MultiSelect, Password } from 'primevue';
import { computed, ref } from 'vue';
import { ref } from 'vue';
import { useI18n } from 'vue-i18n';
import { GroupIdentity, GroupInfo, ensureGroupInfo } from '../../types/network';
import { GroupIdentity, GroupInfo } from '../../types/network';
const props = defineProps<{
groupNames?: string[]
@@ -18,17 +18,6 @@ const editingGroupIndex = ref(-1)
const showGroupDialog = ref(false)
const oldGroupName = ref('')
function groupInfo() {
return ensureGroupInfo(group.value)
}
const members = computed({
get: () => groupInfo().members,
set: value => {
groupInfo().members = value
},
})
function addGroup() {
editingGroupIndex.value = -1
editingGroup.value = {
@@ -41,13 +30,13 @@ function addGroup() {
function editGroup(index: number) {
editingGroupIndex.value = index
editingGroup.value = JSON.parse(JSON.stringify(groupInfo().declares[index]))
editingGroup.value = JSON.parse(JSON.stringify(group.value.declares[index]))
oldGroupName.value = editingGroup.value?.group_name || ''
showGroupDialog.value = true
}
function deleteGroup(index: number) {
groupInfo().declares.splice(index, 1)
group.value.declares.splice(index, 1)
}
function saveGroup() {
@@ -55,15 +44,15 @@ function saveGroup() {
const newName = editingGroup.value.group_name
if (editingGroupIndex.value === -1) {
groupInfo().declares.push(editingGroup.value)
group.value.declares.push(editingGroup.value)
} else {
if (oldGroupName.value && oldGroupName.value !== newName) {
// Sync in members
groupInfo().members = groupInfo().members.map(m => m === oldGroupName.value ? newName : m)
group.value.members = group.value.members.map(m => m === oldGroupName.value ? newName : m)
// Notify parent to sync in rules
emit('rename-group', { oldName: oldGroupName.value, newName })
}
groupInfo().declares[editingGroupIndex.value] = editingGroup.value
group.value.declares[editingGroupIndex.value] = editingGroup.value
}
showGroupDialog.value = false
}
@@ -81,7 +70,7 @@ function saveGroup() {
<Button icon="pi pi-plus" :label="t('web.common.add')" severity="success" @click="addGroup" />
</div>
<DataTable :value="groupInfo().declares" responsiveLayout="scroll">
<DataTable :value="group.declares" responsiveLayout="scroll">
<Column field="group_name" :header="t('acl.group.name')" />
<Column field="group_secret" :header="t('acl.group.secret')">
<template #body="{ data }">
@@ -101,7 +90,7 @@ function saveGroup() {
<div class="flex flex-col gap-2">
<label class="font-bold text-lg">{{ t('acl.group.members') }}</label>
<MultiSelect v-model="members" :options="props.groupNames" multiple fluid filter
<MultiSelect v-model="group.members" :options="props.groupNames" multiple fluid filter
:placeholder="t('acl.group.members')" />
</div>
@@ -2,7 +2,7 @@
import { Button, Menu, Tab, TabList, TabPanel, TabPanels, Tabs } from 'primevue'
import { computed, ref } from 'vue'
import { useI18n } from 'vue-i18n'
import { Acl, AclAction, AclChainType, ensureAclV1 } from '../../types/network'
import { Acl, AclAction, AclChainType } from '../../types/network'
import AclChainEditor from './AclChainEditor.vue'
import AclGroupEditor from './AclGroupEditor.vue'
@@ -12,7 +12,6 @@ const { t } = useI18n()
const activeTab = ref(0)
const menu = ref()
const aclV1 = computed(() => ensureAclV1(acl.value))
const addMenuModel = ref([
{ label: () => t('acl.inbound'), command: () => addChain(AclChainType.Inbound) },
@@ -21,6 +20,10 @@ const addMenuModel = ref([
])
function addChain(type: AclChainType) {
if (!acl.value.acl_v1) {
acl.value.acl_v1 = { chains: [], group: { declares: [], members: [] } }
}
let defaultName = ''
switch (type) {
case AclChainType.Inbound: defaultName = 'Inbound'; break;
@@ -28,7 +31,7 @@ function addChain(type: AclChainType) {
case AclChainType.Forward: defaultName = 'Forward'; break;
}
aclV1.value.chains.push({
acl.value.acl_v1.chains.push({
name: defaultName,
chain_type: type,
description: '',
@@ -37,20 +40,21 @@ function addChain(type: AclChainType) {
default_action: AclAction.Allow
})
activeTab.value = aclV1.value.chains.length - 1
activeTab.value = acl.value.acl_v1.chains.length - 1
}
function removeChain(index: number) {
if (confirm(t('acl.delete_chain_confirm'))) {
aclV1.value.chains.splice(index, 1)
if (activeTab.value >= aclV1.value.chains.length) {
activeTab.value = Math.max(0, aclV1.value.chains.length)
acl.value.acl_v1?.chains.splice(index, 1)
if (activeTab.value >= (acl.value.acl_v1?.chains.length || 0)) {
activeTab.value = Math.max(0, (acl.value.acl_v1?.chains.length || 0))
}
}
}
function handleRenameGroup({ oldName, newName }: { oldName: string, newName: string }) {
aclV1.value.chains.forEach(chain => {
if (!acl.value.acl_v1) return
acl.value.acl_v1.chains.forEach(chain => {
chain.rules.forEach(rule => {
rule.source_groups = rule.source_groups.map(g => g === oldName ? newName : g)
rule.destination_groups = rule.destination_groups.map(g => g === oldName ? newName : g)
@@ -59,11 +63,11 @@ function handleRenameGroup({ oldName, newName }: { oldName: string, newName: str
}
const groupNames = computed(() => {
return aclV1.value.group?.declares.map(g => g.group_name) || []
return acl.value.acl_v1?.group?.declares.map(g => g.group_name) || []
})
const tabs = computed(() => {
const chains = aclV1.value.chains
const chains = acl.value.acl_v1?.chains || []
const result: { type: string, label: string, index: number }[] = []
if (chains.length === 0) {
@@ -120,13 +124,24 @@ const tabs = computed(() => {
</div>
<!-- Rule Chains -->
<div v-if="tab.type === 'chain' && aclV1.chains[tab.index]" class="py-4">
<AclChainEditor v-model="aclV1.chains[tab.index]" :group-names="groupNames" />
<div v-if="tab.type === 'chain' && acl.acl_v1 && acl.acl_v1.chains[tab.index]" class="py-4">
<AclChainEditor v-model="acl.acl_v1.chains[tab.index]" :group-names="groupNames" />
</div>
<!-- Group Management -->
<div v-if="tab.type === 'groups'" class="py-4">
<AclGroupEditor v-model="aclV1.group" :group-names="groupNames" @rename-group="handleRenameGroup" />
<template v-if="acl.acl_v1">
<AclGroupEditor v-if="acl.acl_v1.group" v-model="acl.acl_v1.group" :group-names="groupNames"
@rename-group="handleRenameGroup" />
<div v-else class="flex justify-center p-4">
<Button :label="t('web.common.add') + ' ' + t('acl.groups')"
@click="acl.acl_v1.group = { declares: [], members: [] }" />
</div>
</template>
<div v-else class="flex justify-center p-4">
<Button :label="t('acl.enabled')"
@click="acl.acl_v1 = { chains: [], group: { declares: [], members: [] } }" />
</div>
</div>
</TabPanel>
</TabPanels>
@@ -1,8 +1,8 @@
<script setup lang="ts">
import { AutoComplete, Button, Checkbox, Dialog, InputNumber, InputText, MultiSelect, Panel, SelectButton, ToggleButton } from 'primevue';
import { computed, ref, watch } from 'vue';
import { computed, ref } from 'vue';
import { useI18n } from 'vue-i18n';
import { AclAction, AclProtocol, AclRule, ensureAclRuleLists } from '../../types/network';
import { AclAction, AclProtocol, AclRule } from '../../types/network';
const props = defineProps<{
visible: boolean
@@ -32,8 +32,6 @@ const showPorts = computed(() => {
return rule.value.protocol === AclProtocol.TCP || rule.value.protocol === AclProtocol.UDP || rule.value.protocol === AclProtocol.Any
})
watch(() => rule.value, ensureAclRuleLists, { immediate: true })
function close() {
emit('update:visible', false)
}
@@ -341,8 +341,6 @@ web:
import_config: 导入配置
create_new: 创建新网络
network_status: 网络状态
loading_network_status: 正在加载网络状态
network_info_unavailable: 网络状态不可用
network_configuration: 网络配置
loading_network_configuration: 加载网络配置
no_network_selected: 未选择网络
@@ -341,8 +341,6 @@ web:
import_config: Import Config
create_new: Create New Network
network_status: Network Status
loading_network_status: Loading Network Status
network_info_unavailable: Network status is unavailable
network_configuration: Network Configuration
loading_network_configuration: Loading Network Configuration
no_network_selected: No Network Selected
@@ -1,82 +0,0 @@
import type { PeerRoutePair } from '../types/network'
export function numericValue(value: unknown): number | undefined {
if (typeof value === 'number')
return Number.isFinite(value) ? value : undefined
if (typeof value !== 'string' || value.trim() === '')
return undefined
const parsed = Number(value)
return Number.isFinite(parsed) ? parsed : undefined
}
export function peerConns(info: PeerRoutePair) {
return info.peer?.conns || []
}
function defaultConnId(info: PeerRoutePair) {
const defaultConn = info.peer?.default_conn_id
if (!defaultConn)
return undefined
const part1 = defaultConn.part1 ?? 0
const part2 = defaultConn.part2 ?? 0
const part3 = defaultConn.part3 ?? 0
const part4 = defaultConn.part4 ?? 0
if (part1 === 0 && part2 === 0 && part3 === 0 && part4 === 0)
return undefined
const toHex = (value: number) => value.toString(16).padStart(8, '0')
const part1Hex = toHex(part1)
const part2Hex = toHex(part2)
const part3Hex = toHex(part3)
const part4Hex = toHex(part4)
return `${part1Hex}-${part2Hex.slice(0, 4)}-${part2Hex.slice(4, 8)}-${part3Hex.slice(0, 4)}-${part3Hex.slice(4, 8)}${part4Hex}`
}
function defaultConnFirst(info: PeerRoutePair) {
const conns = peerConns(info)
const connId = defaultConnId(info)
if (!connId)
return conns
const defaultConn = conns.find(conn => conn.conn_id === connId)
return defaultConn ? [defaultConn, ...conns.filter(conn => conn !== defaultConn)] : conns
}
export function latencyMs(info: PeerRoutePair) {
const connId = defaultConnId(info)
let minLatencyUs: number | undefined
for (const conn of peerConns(info)) {
if (!conn.stats)
continue
const latencyUs = numericValue(conn.stats.latency_us)
if (latencyUs === undefined)
continue
if (connId === conn.conn_id)
return `${Math.ceil(latencyUs / 1000)}ms`
minLatencyUs = Math.min(minLatencyUs ?? latencyUs, latencyUs)
}
if (minLatencyUs === undefined)
return ''
return `${Math.ceil(minLatencyUs / 1000)}ms`
}
export function lossRate(info: PeerRoutePair) {
for (const conn of defaultConnFirst(info)) {
const loss = numericValue(conn.loss_rate)
if (loss === undefined)
continue
return `${Math.round(loss * 100)}%`
}
return ''
}
+17 -27
View File
@@ -1,30 +1,24 @@
import { IPv4, IPv6 } from 'ip-num/IPNumber'
import { Ipv4Addr, Ipv4Inet, Ipv6Addr } from '../types/network'
export function ipv4ToString(ip: Ipv4Addr | null | undefined) {
if (!ip) {
return ''
}
return IPv4.fromNumber(ip.addr ?? 0).toString()
export function ipv4ToString(ip: Ipv4Addr) {
return IPv4.fromNumber(ip.addr).toString()
}
export function ipv4InetToString(ip: Ipv4Inet | undefined) {
if (ip?.address === undefined) {
return 'undefined'
}
return `${ipv4ToString(ip.address)}/${ip.network_length ?? 0}`
return `${ipv4ToString(ip.address)}/${ip.network_length}`
}
export function ipv6ToString(ip: Ipv6Addr | null | undefined) {
if (!ip) {
return ''
}
export function ipv6ToString(ip: Ipv6Addr) {
return IPv6.fromBigInt(
(BigInt(ip.part1 ?? 0) << BigInt(96))
+ (BigInt(ip.part2 ?? 0) << BigInt(64))
+ (BigInt(ip.part3 ?? 0) << BigInt(32))
+ BigInt(ip.part4 ?? 0),
).toString()
(BigInt(ip.part1) << BigInt(96))
+ (BigInt(ip.part2) << BigInt(64))
+ (BigInt(ip.part3) << BigInt(32))
+ BigInt(ip.part4),
)
}
function toHexString(uint64: bigint, padding = 9): string {
@@ -49,17 +43,14 @@ function uint32ToUuid(part1: number, part2: number, part3: number, part4: number
}
export interface UUID {
part1?: number;
part2?: number;
part3?: number;
part4?: number;
part1: number;
part2: number;
part3: number;
part4: number;
}
export function UuidToStr(uuid: UUID | null | undefined): string {
if (!uuid) {
return '';
}
return uint32ToUuid(uuid.part1 ?? 0, uuid.part2 ?? 0, uuid.part3 ?? 0, uuid.part4 ?? 0);
export function UuidToStr(uuid: UUID): string {
return uint32ToUuid(uuid.part1, uuid.part2, uuid.part3, uuid.part4);
}
export interface Location {
@@ -80,12 +71,11 @@ export interface DeviceInfo {
}
export function buildDeviceInfo(device: any): DeviceInfo {
const runningInstances = device.info?.running_network_instances ?? [];
let dev_info: DeviceInfo = {
hostname: device.info?.hostname,
public_ip: device.client_url,
running_network_instances: runningInstances.map((instance: any) => UuidToStr(instance)),
running_network_count: runningInstances.length,
running_network_instances: device.info?.running_network_instances.map((instance: any) => UuidToStr(instance)),
running_network_count: device.info?.running_network_instances.length,
report_time: device.info?.report_time,
easytier_version: device.info?.easytier_version,
machine_id: UuidToStr(device.info?.machine_id),
+180 -239
View File
@@ -1,80 +1,166 @@
import { v4 as uuidv4 } from 'uuid'
import {
NetworkConfig as NetworkConfigPb,
NetworkingMethod,
type NetworkPeerConfig,
type NetworkConfig as ProtoNetworkConfig,
type PortForwardConfig,
} from '../generated/proto/api_manage'
import {
Action as AclAction,
ChainType as AclChainType,
Protocol as AclProtocol,
type Acl,
type AclV1,
type Chain as AclChain,
type GroupIdentity,
type GroupInfo,
type Rule as AclRule,
} from '../generated/proto/acl'
import {
CompressionAlgoPb,
NatType,
type PeerFeatureFlag,
type SecureModeConfig,
} from '../generated/proto/common'
import { prepareNetworkConfigForProtoJson } from './networkCompat'
export { AclAction, AclChainType, AclProtocol, CompressionAlgoPb, NatType, NetworkingMethod }
export type { Acl, AclChain, AclRule, AclV1, GroupIdentity, GroupInfo, NetworkPeerConfig, PeerFeatureFlag, PortForwardConfig, SecureModeConfig }
export enum NetworkingMethod {
PublicServer = 0,
Manual = 1,
Standalone = 2,
}
export type NetworkConfig = Omit<
ProtoNetworkConfig,
'instance_id' | 'instance_recv_bps_limit' | 'mtu' | 'networking_method'
> & {
export interface SecureModeConfig {
enabled: boolean
// Keep protocol compatibility with backend/import-export flows even though the GUI
// does not render secure-mode or credential inputs.
local_private_key?: string
local_public_key?: string
}
export enum AclProtocol {
Unspecified = 0,
TCP = 1,
UDP = 2,
ICMP = 3,
ICMPv6 = 4,
Any = 5,
}
export enum AclAction {
Noop = 0,
Allow = 1,
Drop = 2,
}
export enum AclChainType {
UnspecifiedChain = 0,
Inbound = 1,
Outbound = 2,
Forward = 3,
}
export interface AclRule {
name: string
description: string
priority: number
enabled: boolean
protocol: AclProtocol
ports: string[]
source_ips: string[]
destination_ips: string[]
source_ports: string[]
action: AclAction
rate_limit: number
burst_limit: number
stateful: boolean
source_groups: string[]
destination_groups: string[]
}
export interface AclChain {
name: string
chain_type: AclChainType
description: string
enabled: boolean
rules: AclRule[]
default_action: AclAction
}
export interface GroupIdentity {
group_name: string
group_secret: string
}
export interface GroupInfo {
declares: GroupIdentity[]
members: string[]
}
export interface AclV1 {
chains: AclChain[]
group?: GroupInfo
}
export interface Acl {
acl_v1?: AclV1
}
export interface NetworkConfig {
instance_id: string
mtu: number | null
instance_recv_bps_limit: number | string | null
networking_method: NetworkingMethod | string
}
export type NormalizedAclV1 = AclV1 & {
group: GroupInfo
}
dhcp: boolean
virtual_ipv4: string
network_length: number
hostname?: string
network_name: string
network_secret?: string
credential_file?: string
secure_mode?: SecureModeConfig
const UINT64_MAX = (1n << 64n) - 1n
networking_method: NetworkingMethod
interface NetworkingConfigFields {
public_server_url: string
peer_urls: string[]
peers?: NetworkPeerConfig[]
public_server_url?: string
networking_method?: NetworkingMethod | string
}
interface NetworkingMethodOptions {
fillPeerUrlsFromPeers?: boolean
}
proxy_cidrs: string[]
function emptyGroupInfo(): GroupInfo {
return {
declares: [],
members: [],
}
}
enable_vpn_portal: boolean
vpn_portal_listen_port: number
vpn_portal_client_network_addr: string
vpn_portal_client_network_len: number
function emptyAcl(): Acl {
return {
acl_v1: {
group: emptyGroupInfo(),
chains: [],
},
}
advanced_settings: boolean
listener_urls: string[]
latency_first: boolean
dev_name: string
use_smoltcp?: boolean
disable_ipv6?: boolean
ipv6_public_addr_auto?: boolean
enable_kcp_proxy?: boolean
disable_kcp_input?: boolean
enable_quic_proxy?: boolean
disable_quic_input?: boolean
disable_p2p?: boolean
p2p_only?: boolean
lazy_p2p?: boolean
bind_device?: boolean
no_tun?: boolean
enable_exit_node?: boolean
relay_all_peer_rpc?: boolean
need_p2p?: boolean
multi_thread?: boolean
proxy_forward_by_system?: boolean
disable_encryption?: boolean
disable_tcp_hole_punching?: boolean
disable_udp_hole_punching?: boolean
disable_upnp?: boolean
enable_udp_broadcast_relay?: boolean
disable_sym_hole_punching?: boolean
enable_relay_network_whitelist?: boolean
relay_network_whitelist: string[]
enable_manual_routes: boolean
routes: string[]
exit_nodes: string[]
enable_socks5?: boolean
socks5_port: number
mtu: number | null
instance_recv_bps_limit: number | null
mapped_listeners: string[]
enable_magic_dns?: boolean
enable_private_mode?: boolean
port_forwards: PortForwardConfig[]
acl?: Acl
}
export function DEFAULT_NETWORK_CONFIG(): NetworkConfig {
return {
...NetworkConfigPb.create(),
instance_id: uuidv4(),
dhcp: true,
@@ -141,7 +227,15 @@ export function DEFAULT_NETWORK_CONFIG(): NetworkConfig {
enable_magic_dns: false,
enable_private_mode: false,
port_forwards: [],
acl: emptyAcl(),
acl: {
acl_v1: {
group: {
declares: [],
members: [],
},
chains: [],
},
},
}
}
@@ -149,185 +243,33 @@ function cleanPeerUrls(urls: string[] | undefined): string[] {
return (urls ?? []).map((url) => url.trim()).filter((url) => url.length > 0)
}
function cleanNetworkPeers(peers: NetworkPeerConfig[] | undefined): NetworkPeerConfig[] {
return (peers ?? [])
.map((peer) => ({
...peer,
uri: peer.uri.trim(),
}))
.filter((peer) => peer.uri.length > 0)
}
function peersFromUrls(urls: string[], existingPeers: NetworkPeerConfig[]): NetworkPeerConfig[] {
const peersByUri = new Map<string, NetworkPeerConfig>()
for (const peer of existingPeers) {
if (!peersByUri.has(peer.uri)) {
peersByUri.set(peer.uri, peer)
}
export function normalizeNetworkConfig(config: NetworkConfig): NetworkConfig {
const normalized: NetworkConfig = {
...config,
peer_urls: cleanPeerUrls(config.peer_urls),
}
return urls.map((uri) => ({
...(peersByUri.get(uri) ?? {}),
uri,
}))
}
const publicServerUrl = normalized.public_server_url?.trim() ?? ''
export function ensureAclRuleLists(rule: AclRule): AclRule {
rule.ports ??= []
rule.source_ips ??= []
rule.destination_ips ??= []
rule.source_ports ??= []
rule.source_groups ??= []
rule.destination_groups ??= []
return rule
}
export function ensureAclChain(chain: AclChain): AclChain {
chain.rules ??= []
chain.rules.forEach(ensureAclRuleLists)
return chain
}
export function ensureGroupInfo(group: GroupInfo): GroupInfo {
group.declares ??= []
group.members ??= []
return group
}
export function ensureAclV1(acl: Acl): NormalizedAclV1 {
acl.acl_v1 ??= { chains: [], group: emptyGroupInfo() }
acl.acl_v1.chains ??= []
acl.acl_v1.chains.forEach(ensureAclChain)
acl.acl_v1.group = ensureGroupInfo(acl.acl_v1.group ?? emptyGroupInfo())
return acl.acl_v1 as NormalizedAclV1
}
function normalizeAcl(acl: Acl | undefined): Acl {
const source = acl ?? emptyAcl()
const aclV1 = source.acl_v1 ?? { chains: [], group: emptyGroupInfo() }
return {
...source,
acl_v1: {
...aclV1,
chains: (aclV1.chains ?? []).map((chain) => ({
...chain,
rules: (chain.rules ?? []).map((rule) => ({ ...ensureAclRuleLists({ ...rule }) })),
})),
group: ensureGroupInfo({
...(aclV1.group ?? emptyGroupInfo()),
declares: aclV1.group?.declares ?? [],
members: aclV1.group?.members ?? [],
}),
},
}
}
function isGroupInfoEmpty(group: GroupInfo | undefined): boolean {
return (group?.declares?.length ?? 0) === 0 && (group?.members?.length ?? 0) === 0
}
function isAclEmpty(acl: Acl | undefined): boolean {
const aclV1 = acl?.acl_v1
return !aclV1 || ((aclV1.chains?.length ?? 0) === 0 && isGroupInfoEmpty(aclV1.group))
}
function normalizeUint64ForInput(v: bigint | number | string | null | undefined): number | string | null {
if (v == null) return null
try {
const n = typeof v === 'bigint' ? v : BigInt(v)
if (n === 0n || n > UINT64_MAX) return null
return n <= BigInt(Number.MAX_SAFE_INTEGER) ? Number(n) : n.toString()
} catch {
return null
}
}
function normalizeNumberForInput(v: number | string | null | undefined): number | null {
if (v == null) return null
const n = Number(v)
return Number.isFinite(n) ? n : null
}
function toBackendUint64(v: number | bigint | string | null | undefined): bigint | undefined {
if (v == null || v === '') return undefined
try {
const n = typeof v === 'bigint' ? v : BigInt(v)
return n > 0n && n <= UINT64_MAX ? n : undefined
} catch {
return undefined
}
}
function applyNetworkingMethod(
config: NetworkingConfigFields,
options: NetworkingMethodOptions = {},
): void {
const existingPeers = cleanNetworkPeers(config.peers)
config.peer_urls = cleanPeerUrls(config.peer_urls)
if (options.fillPeerUrlsFromPeers && config.peer_urls.length === 0 && existingPeers.length > 0) {
config.peer_urls = existingPeers.map((peer) => peer.uri)
}
const publicServerUrl = config.public_server_url?.trim() ?? ''
const networkingMethod = config.networking_method ?? NetworkingMethod.Manual
switch (networkingMethod) {
switch (normalized.networking_method) {
case NetworkingMethod.PublicServer:
config.peer_urls = publicServerUrl
? [publicServerUrl]
: (options.fillPeerUrlsFromPeers ? existingPeers.map((peer) => peer.uri) : [])
normalized.peer_urls = publicServerUrl ? [publicServerUrl] : []
break
case NetworkingMethod.Manual:
break
case NetworkingMethod.Standalone:
default:
config.peer_urls = []
normalized.peer_urls = []
break
}
config.networking_method = NetworkingMethod.Manual
config.public_server_url = ''
config.peers = peersFromUrls(config.peer_urls, existingPeers)
}
export function normalizeNetworkConfig(config: NetworkConfig): NetworkConfig {
const normalized = NetworkConfigPb.fromJson(prepareNetworkConfigForProtoJson(config) as any, {
ignoreUnknownFields: true,
}) as unknown as NetworkConfig
applyNetworkingMethod(normalized, { fillPeerUrlsFromPeers: true })
normalized.mtu = normalizeNumberForInput(normalized.mtu)
normalized.instance_recv_bps_limit = normalizeUint64ForInput(
normalized.instance_recv_bps_limit as any,
)
normalized.proxy_cidrs ??= []
normalized.listener_urls ??= []
normalized.relay_network_whitelist ??= []
normalized.routes ??= []
normalized.exit_nodes ??= []
normalized.mapped_listeners ??= []
normalized.port_forwards ??= []
normalized.acl = config.acl === undefined ? undefined : normalizeAcl(normalized.acl)
normalized.networking_method = NetworkingMethod.Manual
normalized.public_server_url = ''
return normalized
}
export function toBackendNetworkConfig(config: NetworkConfig): NetworkConfig {
const backend = NetworkConfigPb.fromJson(prepareNetworkConfigForProtoJson(config) as any, {
ignoreUnknownFields: true,
})
applyNetworkingMethod(backend)
backend.mtu = normalizeNumberForInput(config.mtu) ?? undefined
backend.instance_recv_bps_limit = toBackendUint64(config.instance_recv_bps_limit)
if (config.acl === undefined || isAclEmpty(config.acl)) {
backend.acl = undefined
}
return NetworkConfigPb.toJson(backend, {
useProtoFieldName: true,
}) as unknown as NetworkConfig
return normalizeNetworkConfig(config)
}
export interface NetworkInstance {
@@ -412,7 +354,6 @@ export interface Route {
proxy_cidrs: string[]
hostname: string
stun_info?: StunInfo
feature_flag?: PeerFeatureFlag
inst_id: string
version: string
}
@@ -420,7 +361,6 @@ export interface Route {
export interface PeerInfo {
peer_id: number
conns: PeerConnInfo[]
default_conn_id?: CommonUuid
}
export interface PeerConnInfo {
@@ -431,7 +371,7 @@ export interface PeerConnInfo {
features: string[]
tunnel?: TunnelInfo
stats?: PeerConnStats
loss_rate?: number | string
loss_rate: number
}
export interface PeerRoutePair {
@@ -450,18 +390,19 @@ export interface TunnelInfo {
}
export interface PeerConnStats {
rx_bytes: number | string
tx_bytes: number | string
rx_packets: number | string
tx_packets: number | string
latency_us: number | string
rx_bytes: number
tx_bytes: number
rx_packets: number
tx_packets: number
latency_us: number
}
export interface CommonUuid {
part1?: number
part2?: number
part3?: number
part4?: number
export interface PortForwardConfig {
bind_ip: string,
bind_port: number,
dst_ip: string,
dst_port: number,
proto: string
}
// 添加新行
@@ -1,85 +0,0 @@
import {
Action as AclAction,
ChainType as AclChainType,
Protocol as AclProtocol,
} from '../generated/proto/acl'
import type { NetworkConfig } from './network'
const UINT64_MAX = (1n << 64n) - 1n
type JsonRecord = Record<string, unknown>
export function prepareNetworkConfigForProtoJson(config: NetworkConfig): NetworkConfig {
const prepared = dropUnsupportedJsonValues(applyLegacyAclDefaults(config)) as NetworkConfig
normalizeLegacyOptionalUint64(prepared as JsonRecord, 'instance_recv_bps_limit')
return prepared
}
function applyLegacyAclDefaults(config: NetworkConfig): NetworkConfig {
const acl = config.acl
const aclV1 = acl?.acl_v1
if (!Array.isArray(aclV1?.chains)) return config
return {
...config,
acl: {
...acl,
acl_v1: {
...aclV1,
chains: aclV1.chains.map((chain) => ({
...chain,
chain_type: chain.chain_type ?? AclChainType.UnspecifiedChain,
default_action: chain.default_action ?? AclAction.Allow,
rules: (chain.rules ?? []).map((rule) => ({
...rule,
protocol: rule.protocol ?? AclProtocol.Any,
action: rule.action ?? AclAction.Allow,
})),
})),
},
},
}
}
function dropUnsupportedJsonValues(value: unknown): unknown {
if (value === undefined) return undefined
if (typeof value === 'number' && !Number.isFinite(value)) return undefined
if (Array.isArray(value)) {
return value.map(dropUnsupportedJsonValues).filter((v) => v !== undefined)
}
if (isJsonRecord(value)) {
return Object.fromEntries(
Object.entries(value)
.map(([k, v]) => [k, dropUnsupportedJsonValues(v)])
.filter(([, v]) => v !== undefined),
)
}
return value
}
function isJsonRecord(value: unknown): value is JsonRecord {
return typeof value === 'object' && value !== null
}
function normalizeLegacyOptionalUint64(obj: JsonRecord, key: string): void {
const value = obj[key]
if (typeof value !== 'string') return
const trimmed = value.trim()
if (!isPositiveUint64String(trimmed)) {
delete obj[key]
return
}
obj[key] = trimmed
}
function isPositiveUint64String(value: string): boolean {
if (!/^\d+$/.test(value)) return false
const n = BigInt(value)
return n > 0n && n <= UINT64_MAX
}
@@ -1,563 +0,0 @@
import { mount, type VueWrapper } from '@vue/test-utils'
import { describe, expect, it, vi } from 'vitest'
import { defineComponent, h, nextTick, reactive } from 'vue'
import Config from '../src/components/Config.vue'
import {
DEFAULT_NETWORK_CONFIG,
toBackendNetworkConfig,
type NetworkConfig,
} from '../src/types/network'
const CONFIG_FLAG_FIELDS = [
'latency_first',
'use_smoltcp',
'disable_ipv6',
'ipv6_public_addr_auto',
'enable_kcp_proxy',
'disable_kcp_input',
'enable_quic_proxy',
'disable_quic_input',
'disable_p2p',
'p2p_only',
'lazy_p2p',
'bind_device',
'no_tun',
'enable_exit_node',
'relay_all_peer_rpc',
'need_p2p',
'multi_thread',
'proxy_forward_by_system',
'disable_encryption',
'disable_tcp_hole_punching',
'disable_udp_hole_punching',
'enable_udp_broadcast_relay',
'disable_upnp',
'disable_sym_hole_punching',
'enable_magic_dns',
'enable_private_mode',
] as const satisfies readonly (keyof NetworkConfig)[]
const CONFIG_CHECKBOX_FIELDS = [
['dhcp', '#virtual_ip_auto'],
...CONFIG_FLAG_FIELDS.map((field) => [field, `#${field}`] as const),
] as const satisfies readonly (readonly [keyof NetworkConfig, string])[]
const CONFIG_TOGGLE_FIELDS = [
'enable_vpn_portal',
'enable_relay_network_whitelist',
'enable_manual_routes',
'enable_socks5',
] as const satisfies readonly (keyof NetworkConfig)[]
const CONFIG_UI_BOOLEAN_FIELDS = [
...CONFIG_CHECKBOX_FIELDS.map(([field]) => field),
...CONFIG_TOGGLE_FIELDS,
] as const satisfies readonly (keyof NetworkConfig)[]
vi.mock('vue-i18n', () => ({
useI18n: () => ({
t: (key: string, values?: unknown[]) => values ? `${key}:${values.join(',')}` : key,
}),
}))
const PassThrough = defineComponent({
name: 'PassThrough',
setup(_, { slots }) {
return () => h('div', slots.default?.())
},
})
const PanelStub = defineComponent({
name: 'Panel',
props: {
header: String,
},
setup(props, { slots }) {
return () => h('section', { 'data-stub': 'panel', 'data-header': props.header }, slots.default?.())
},
})
const DividerStub = defineComponent({
name: 'Divider',
setup() {
return () => h('hr', { 'data-stub': 'divider' })
},
})
function splitList(value: string): string[] {
return value.split(',').map((item) => item.trim()).filter((item) => item.length > 0)
}
const InputTextStub = defineComponent({
name: 'InputText',
props: {
modelValue: [String, Number],
id: String,
disabled: Boolean,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.id,
disabled: props.disabled,
value: props.modelValue ?? '',
'data-stub': 'input-text',
onInput: (event: Event) => emit('update:modelValue', (event.target as HTMLInputElement).value),
})
},
})
const PasswordStub = defineComponent({
name: 'Password',
props: {
modelValue: [String, Number],
id: String,
disabled: Boolean,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.id,
disabled: props.disabled,
type: 'password',
value: props.modelValue ?? '',
'data-stub': 'password',
onInput: (event: Event) => emit('update:modelValue', (event.target as HTMLInputElement).value),
})
},
})
const InputNumberStub = defineComponent({
name: 'InputNumber',
props: {
modelValue: Number,
id: String,
inputId: String,
disabled: Boolean,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.id ?? props.inputId,
disabled: props.disabled,
type: 'number',
value: props.modelValue ?? '',
'data-stub': 'input-number',
onInput: (event: Event) => {
const value = (event.target as HTMLInputElement).value
emit('update:modelValue', value === '' ? null : Number(value))
},
})
},
})
const CheckboxStub = defineComponent({
name: 'Checkbox',
props: {
modelValue: Boolean,
inputId: String,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.inputId,
checked: props.modelValue,
type: 'checkbox',
'data-stub': 'checkbox',
onChange: (event: Event) => emit('update:modelValue', (event.target as HTMLInputElement).checked),
})
},
})
const ToggleButtonStub = defineComponent({
name: 'ToggleButton',
props: {
modelValue: Boolean,
onIcon: String,
offIcon: String,
onLabel: String,
offLabel: String,
},
emits: ['update:modelValue'],
setup(props, { emit }) {
return () => h('button', {
type: 'button',
'aria-pressed': String(Boolean(props.modelValue)),
'data-stub': 'toggle-button',
onClick: () => emit('update:modelValue', !props.modelValue),
}, props.modelValue ? props.onLabel : props.offLabel)
},
})
const AutoCompleteStub = defineComponent({
name: 'AutoComplete',
props: {
modelValue: Array,
id: String,
multiple: Boolean,
},
emits: ['update:modelValue', 'complete'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.id,
value: (props.modelValue ?? []).join(','),
'data-stub': 'auto-complete',
onInput: (event: Event) => emit('update:modelValue', splitList((event.target as HTMLInputElement).value)),
})
},
})
const UrlListInputStub = defineComponent({
name: 'UrlListInput',
props: {
modelValue: Array,
id: String,
addLabel: String,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.id,
value: (props.modelValue ?? []).join(','),
'data-stub': 'url-list-input',
'data-add-label': props.addLabel,
onInput: (event: Event) => emit('update:modelValue', splitList((event.target as HTMLInputElement).value)),
})
},
})
const SelectButtonStub = defineComponent({
name: 'SelectButton',
props: {
modelValue: String,
options: Array,
},
emits: ['update:modelValue'],
setup(props, { emit }) {
return () => h('select', {
value: props.modelValue,
'data-stub': 'select-button',
onChange: (event: Event) => emit('update:modelValue', (event.target as HTMLSelectElement).value),
}, (props.options ?? []).map((option) => h('option', { value: option as string }, option as string)))
},
})
const ButtonStub = defineComponent({
name: 'Button',
props: {
label: String,
icon: String,
disabled: Boolean,
},
emits: ['click'],
setup(props, { slots, emit }) {
return () => h('button', {
type: 'button',
disabled: props.disabled,
'data-label': props.label ?? props.icon,
onClick: (event: MouseEvent) => emit('click', event),
}, slots.default?.() ?? props.label ?? props.icon)
},
})
const DialogStub = defineComponent({
name: 'Dialog',
props: {
visible: Boolean,
},
setup(props, { slots }) {
return () => h('div', { hidden: !props.visible, 'data-stub': 'dialog' }, [
slots.default?.(),
slots.footer?.(),
])
},
})
const AclManagerStub = defineComponent({
name: 'AclManager',
props: {
modelValue: Object,
},
emits: ['update:modelValue'],
setup(props) {
return () => h('pre', { 'data-stub': 'acl-manager' }, JSON.stringify(props.modelValue))
},
})
function makeConfig(): NetworkConfig {
const config = DEFAULT_NETWORK_CONFIG()
return {
...config,
dhcp: false,
virtual_ipv4: '10.1.2.3',
network_length: 24,
network_name: 'mesh-a',
network_secret: 'secret-a',
peer_urls: ['tcp://peer-a:11010', 'udp://peer-b:11010'],
latency_first: true,
use_smoltcp: true,
disable_ipv6: true,
no_tun: true,
hostname: 'host-a',
proxy_cidrs: ['10.10.0.0/16', '172.16.1.0/24'],
enable_vpn_portal: true,
vpn_portal_client_network_addr: '10.144.0.0',
vpn_portal_listen_port: 22023,
listener_urls: ['tcp://0.0.0.0:12010'],
dev_name: 'tun-test',
mtu: 1280,
instance_recv_bps_limit: '9007199254740993',
enable_relay_network_whitelist: true,
relay_network_whitelist: ['network-a'],
enable_manual_routes: true,
routes: ['192.168.0.0/16'],
enable_socks5: true,
socks5_port: 1086,
exit_nodes: ['exit-a'],
mapped_listeners: ['tcp://127.0.0.1:22000'],
port_forwards: [{
proto: 'udp',
bind_ip: '0.0.0.0',
bind_port: 18080,
dst_ip: '10.0.0.2',
dst_port: 8080,
}],
}
}
function mountConfig(config: NetworkConfig = makeConfig()) {
const curNetwork = reactive(config) as NetworkConfig
const wrapper = mount(Config, {
props: {
curNetwork,
hostname: 'host-from-prop',
},
global: {
directives: {
tooltip: () => {},
},
stubs: {
AclManager: AclManagerStub,
AutoComplete: AutoCompleteStub,
Button: ButtonStub,
Checkbox: CheckboxStub,
Dialog: DialogStub,
Divider: DividerStub,
InputGroup: PassThrough,
InputGroupAddon: PassThrough,
InputNumber: InputNumberStub,
InputText: InputTextStub,
Panel: PanelStub,
Password: PasswordStub,
SelectButton: SelectButtonStub,
ToggleButton: ToggleButtonStub,
UrlListInput: UrlListInputStub,
},
},
})
return { curNetwork, wrapper }
}
function input(wrapper: VueWrapper, selector: string): HTMLInputElement {
return wrapper.find(selector).element as HTMLInputElement
}
async function setInput(wrapper: VueWrapper, selector: string, value: string) {
await wrapper.find(selector).setValue(value)
await nextTick()
}
describe('Config.vue network config projection', () => {
it('projects config values into the visible form controls', async () => {
const { curNetwork, wrapper } = mountConfig()
await nextTick()
expect(input(wrapper, '#network_name').value).toBe('mesh-a')
expect(input(wrapper, '#network_secret').value).toBe('secret-a')
expect(input(wrapper, '#virtual_ip').value).toBe('10.1.2.3')
expect(input(wrapper, '#initial_nodes').value).toBe('tcp://peer-a:11010,udp://peer-b:11010')
expect(input(wrapper, '#virtual_ip_auto').checked).toBe(false)
expect(input(wrapper, '#latency_first').checked).toBe(true)
expect(input(wrapper, '#use_smoltcp').checked).toBe(true)
expect(input(wrapper, '#disable_ipv6').checked).toBe(true)
expect(input(wrapper, '#no_tun').checked).toBe(true)
expect(input(wrapper, '#hostname').value).toBe('host-a')
expect(input(wrapper, '#subnet-proxy').value).toBe('10.10.0.0/16,172.16.1.0/24')
expect(input(wrapper, 'input[placeholder="vpn_portal_client_network"]').value).toBe('10.144.0.0')
expect(input(wrapper, '#dev_name').value).toBe('tun-test')
expect(input(wrapper, '#mtu').value).toBe('1280')
expect(input(wrapper, '#instance_recv_bps_limit').value).toBe('9007199254740993')
expect(input(wrapper, '#relay_network_whitelist').value).toBe('network-a')
expect(input(wrapper, '#routes').value).toBe('192.168.0.0/16')
expect(input(wrapper, '#socks5_port').value).toBe('1086')
expect(input(wrapper, '#exit_nodes').value).toBe('exit-a')
expect(input(wrapper, 'input[data-add-label="add_listener_url"]').value).toBe('tcp://0.0.0.0:12010')
expect(input(wrapper, 'input[data-add-label="add_mapped_listener"]').value).toBe('tcp://127.0.0.1:22000')
expect(wrapper.find<HTMLSelectElement>('select[data-stub="select-button"]').element.value).toBe('udp')
expect(input(wrapper, 'input[placeholder="port_forwards_bind_addr"]').value).toBe('0.0.0.0')
expect(input(wrapper, 'input[placeholder="port_forwards_dst_addr"]').value).toBe('10.0.0.2')
expect(wrapper.findComponent(AclManagerStub).props('modelValue')).toStrictEqual(curNetwork.acl)
})
it('projects form edits back into config and backend JSON', async () => {
const { curNetwork, wrapper } = mountConfig()
await nextTick()
await wrapper.find('#virtual_ip_auto').setValue(false)
await setInput(wrapper, '#network_name', 'mesh-edited')
await setInput(wrapper, '#network_secret', 'secret-edited')
await setInput(wrapper, '#virtual_ip', '10.7.7.7')
await setInput(wrapper, '#initial_nodes', ' tcp://peer-x:11010, , udp://peer-y:11010 ')
await wrapper.find('#no_tun').setValue(false)
await wrapper.find('#disable_ipv6').setValue(false)
await setInput(wrapper, '#hostname', 'host-edited')
await setInput(wrapper, '#subnet-proxy', '10.7.0.0/16,172.17.0.0/16')
await setInput(wrapper, 'input[placeholder="vpn_portal_client_network"]', '10.200.0.0')
await setInput(wrapper, 'input[data-add-label="add_listener_url"]', 'tcp://0.0.0.0:13010')
await setInput(wrapper, '#dev_name', 'tun-edited')
await setInput(wrapper, '#mtu', '1260')
await setInput(wrapper, '#instance_recv_bps_limit', '9007199254740993')
await setInput(wrapper, '#relay_network_whitelist', 'network-edited')
await setInput(wrapper, '#routes', '192.168.10.0/24')
await setInput(wrapper, '#socks5_port', '1089')
await setInput(wrapper, '#exit_nodes', 'exit-edited')
await setInput(wrapper, 'input[data-add-label="add_mapped_listener"]', 'tcp://127.0.0.1:23000')
await wrapper.find('select[data-stub="select-button"]').setValue('tcp')
await setInput(wrapper, 'input[placeholder="port_forwards_bind_addr"]', '127.0.0.1')
await setInput(wrapper, 'input[placeholder="port_forwards_dst_addr"]', '10.9.0.2')
const portNumbers = wrapper.findAll<HTMLInputElement>('input#horizontal-buttons')
await portNumbers[1].setValue('19090')
await portNumbers[2].setValue('9090')
expect(curNetwork).toMatchObject({
dhcp: false,
virtual_ipv4: '10.7.7.7',
network_name: 'mesh-edited',
network_secret: 'secret-edited',
peer_urls: ['tcp://peer-x:11010', 'udp://peer-y:11010'],
no_tun: false,
disable_ipv6: false,
hostname: 'host-edited',
proxy_cidrs: ['10.7.0.0/16', '172.17.0.0/16'],
vpn_portal_client_network_addr: '10.200.0.0',
listener_urls: ['tcp://0.0.0.0:13010'],
dev_name: 'tun-edited',
mtu: 1260,
instance_recv_bps_limit: '9007199254740993',
relay_network_whitelist: ['network-edited'],
routes: ['192.168.10.0/24'],
socks5_port: 1089,
exit_nodes: ['exit-edited'],
mapped_listeners: ['tcp://127.0.0.1:23000'],
port_forwards: [{
proto: 'tcp',
bind_ip: '127.0.0.1',
bind_port: 19090,
dst_ip: '10.9.0.2',
dst_port: 9090,
}],
})
const backend = toBackendNetworkConfig(curNetwork)
expect(backend).toMatchObject({
virtual_ipv4: '10.7.7.7',
network_name: 'mesh-edited',
network_secret: 'secret-edited',
peer_urls: ['tcp://peer-x:11010', 'udp://peer-y:11010'],
listener_urls: ['tcp://0.0.0.0:13010'],
mtu: 1260,
instance_recv_bps_limit: '9007199254740993',
port_forwards: [{
proto: 'tcp',
bind_ip: '127.0.0.1',
bind_port: 19090,
dst_ip: '10.9.0.2',
dst_port: 9090,
}],
})
})
it('round-trips every visible boolean config control into backend JSON', async () => {
const config = makeConfig()
const originalFlagValues = new Map(
CONFIG_UI_BOOLEAN_FIELDS.map((field, index) => {
const value = index % 2 === 0
config[field] = value
return [field, value]
}),
)
const { curNetwork, wrapper } = mountConfig(config)
await nextTick()
for (const [field, selector] of CONFIG_CHECKBOX_FIELDS) {
const value = originalFlagValues.get(field)
expect(input(wrapper, selector).checked, `${field} should project into UI`).toBe(value)
await wrapper.find(selector).setValue(!value)
await nextTick()
}
const toggleButtons = wrapper.findAll('button[data-stub="toggle-button"]')
expect(toggleButtons).toHaveLength(CONFIG_TOGGLE_FIELDS.length)
for (const [index, field] of CONFIG_TOGGLE_FIELDS.entries()) {
const value = originalFlagValues.get(field)
expect(toggleButtons[index].attributes('aria-pressed'), `${field} should project into UI`)
.toBe(String(value))
await toggleButtons[index].trigger('click')
await nextTick()
}
const backend = toBackendNetworkConfig(curNetwork) as Record<string, unknown>
for (const [field, value] of originalFlagValues) {
const expectedValue = !value
expect(curNetwork[field], `${field} should update config`).toBe(expectedValue)
expect(backend[field], `${field} should be preserved in backend JSON`).toBe(expectedValue)
}
})
it('keeps uint64 input editable without losing large values', async () => {
const { curNetwork, wrapper } = mountConfig()
await nextTick()
await setInput(wrapper, '#instance_recv_bps_limit', '1234')
expect(curNetwork.instance_recv_bps_limit).toBe(1234)
await setInput(wrapper, '#instance_recv_bps_limit', 'not-a-number')
expect(curNetwork.instance_recv_bps_limit).toBe(1234)
await setInput(wrapper, '#instance_recv_bps_limit', '0')
expect(curNetwork.instance_recv_bps_limit).toBeNull()
expect(input(wrapper, '#instance_recv_bps_limit').value).toBe('')
await setInput(wrapper, '#instance_recv_bps_limit', '9007199254740993')
expect(curNetwork.instance_recv_bps_limit).toBe('9007199254740993')
await setInput(wrapper, '#instance_recv_bps_limit', '18446744073709551616')
expect(curNetwork.instance_recv_bps_limit).toBe('9007199254740993')
await setInput(wrapper, '#instance_recv_bps_limit', '')
expect(curNetwork.instance_recv_bps_limit).toBeNull()
})
it('emits runNetwork with the current projected config', async () => {
const { curNetwork, wrapper } = mountConfig()
await nextTick()
await setInput(wrapper, '#network_name', 'mesh-running')
await wrapper.find('button[data-label="run_network"]').trigger('click')
expect(wrapper.emitted('runNetwork')?.[0]).toEqual([curNetwork])
expect((wrapper.emitted('runNetwork')?.[0][0] as NetworkConfig).network_name).toBe('mesh-running')
})
})
@@ -1,228 +0,0 @@
import { flushPromises, mount } from '@vue/test-utils'
import { describe, expect, it, vi } from 'vitest'
import { nextTick } from 'vue'
import RemoteManagement from '../src/components/RemoteManagement.vue'
import {
DEFAULT_NETWORK_CONFIG,
type NetworkConfig,
} from '../src/types/network'
const BOOLEAN_CONFIG_FIELDS = [
'dhcp',
'enable_vpn_portal',
'advanced_settings',
'latency_first',
'use_smoltcp',
'disable_ipv6',
'enable_kcp_proxy',
'disable_kcp_input',
'disable_p2p',
'bind_device',
'no_tun',
'enable_exit_node',
'relay_all_peer_rpc',
'multi_thread',
'enable_relay_network_whitelist',
'enable_manual_routes',
'proxy_forward_by_system',
'disable_encryption',
'enable_socks5',
'disable_udp_hole_punching',
'enable_magic_dns',
'enable_private_mode',
'enable_quic_proxy',
'disable_quic_input',
'disable_sym_hole_punching',
'p2p_only',
'lazy_p2p',
'need_p2p',
'disable_upnp',
'ipv6_public_addr_provider',
'ipv6_public_addr_auto',
'disable_relay_data',
'enable_udp_broadcast_relay',
'disable_tcp_hole_punching',
] as const satisfies readonly (keyof NetworkConfig)[]
vi.mock('vue-i18n', () => ({
useI18n: () => ({
t: (key: string) => key,
}),
}))
vi.mock('primevue', async () => {
const { defineComponent, h } = await import('vue')
const PassThrough = defineComponent({
name: 'PassThrough',
props: {
label: String,
value: String,
},
setup(props, { slots }) {
return () => h('div', {
'data-label': props.label,
'data-value': props.value,
'data-stub': 'pass-through',
}, slots.default?.())
},
})
const ButtonStub = defineComponent({
name: 'Button',
props: {
label: String,
icon: String,
disabled: Boolean,
},
emits: ['click'],
setup(props, { slots, emit }) {
return () => h('button', {
type: 'button',
disabled: props.disabled,
'data-label': props.label ?? props.icon,
onClick: (event: MouseEvent) => emit('click', event),
}, slots.default?.() ?? props.label ?? props.icon)
},
})
const SelectStub = defineComponent({
name: 'Select',
props: {
modelValue: Object,
options: Array,
},
emits: ['update:modelValue'],
setup(props, { slots }) {
return () => h('div', { 'data-stub': 'select' }, [
slots.value?.({ value: props.modelValue, placeholder: '' }),
])
},
})
const MenuStub = defineComponent({
name: 'Menu',
setup(_, { expose }) {
expose({ toggle: vi.fn() })
return () => h('div', { 'data-stub': 'menu' })
},
})
return {
Button: ButtonStub,
ConfirmPopup: PassThrough,
Divider: PassThrough,
IftaLabel: PassThrough,
Menu: MenuStub,
Message: PassThrough,
Select: SelectStub,
Tag: PassThrough,
useConfirm: () => ({ require: vi.fn() }),
useToast: () => ({ add: vi.fn() }),
}
})
const INSTANCE_ID = '00000000-0000-0000-0000-000000000001'
const INSTANCE_UUID = {
part1: 0,
part2: 0,
part3: 0,
part4: 1,
}
function makeFlagConfig(): NetworkConfig {
const config = {
...DEFAULT_NETWORK_CONFIG(),
instance_id: INSTANCE_ID,
network_name: 'mesh-save',
}
BOOLEAN_CONFIG_FIELDS.forEach((field, index) => {
config[field] = index % 2 === 0
})
return config
}
function cloneConfig(config: NetworkConfig): NetworkConfig {
return JSON.parse(JSON.stringify(config)) as NetworkConfig
}
function snapshotBooleanConfigFields(config: NetworkConfig): Record<string, unknown> {
return Object.fromEntries(
BOOLEAN_CONFIG_FIELDS.map((field) => [field, config[field]]),
)
}
async function settleRemoteManagement() {
for (let i = 0; i < 3; i++) {
await new Promise((resolve) => setTimeout(resolve, 0))
await flushPromises()
await nextTick()
}
}
describe('RemoteManagement config save', () => {
it('saves the current network config without dropping boolean fields', async () => {
const config = makeFlagConfig()
const expectedFlags = snapshotBooleanConfigFields(config)
const api = {
delete_network: vi.fn(),
generate_config: vi.fn(),
get_network_config: vi.fn(async () => cloneConfig(config)),
get_network_info: vi.fn(),
get_network_metas: vi.fn(async (instanceIds: string[]) => ({
metas: Object.fromEntries(instanceIds.map((id) => [id, {
config_permission: 0xffffffff,
inst_id: INSTANCE_UUID,
instance_name: 'mesh-save',
network_name: 'mesh-save',
source: 2,
}])),
})),
list_network_instance_ids: vi.fn(async () => ({
disabled_inst_ids: [INSTANCE_UUID],
running_inst_ids: [],
})),
parse_config: vi.fn(),
run_network: vi.fn(),
save_config: vi.fn(async () => undefined),
update_network_instance_state: vi.fn(),
validate_config: vi.fn(),
}
const wrapper = mount(RemoteManagement, {
props: {
api,
instanceId: INSTANCE_ID,
},
global: {
stubs: {
Config: true,
ConfigEditDialog: true,
Status: true,
},
},
})
try {
await settleRemoteManagement()
const saveButton = wrapper.find('button[data-label="web.device_management.save_config"]')
expect(saveButton.exists()).toBe(true)
expect(saveButton.attributes('disabled')).toBeUndefined()
await saveButton.trigger('click')
await flushPromises()
expect(api.save_config).toHaveBeenCalledOnce()
const savedConfig = api.save_config.mock.calls[0][0] as NetworkConfig
for (const field of BOOLEAN_CONFIG_FIELDS) {
expect(savedConfig[field], `${field} should be saved`).toBe(expectedFlags[field])
}
} finally {
wrapper.unmount()
}
})
})
-9
View File
@@ -1,9 +0,0 @@
import { vi } from 'vitest'
class ResizeObserverStub {
observe() {}
unobserve() {}
disconnect() {}
}
vi.stubGlobal('ResizeObserver', ResizeObserverStub)
@@ -1,84 +0,0 @@
import { describe, expect, it } from 'vitest'
import { latencyMs, lossRate } from '../src/modules/statusDisplay'
import { ipv4ToString, ipv6ToString } from '../src/modules/utils'
function peerRoutePair(conns: any[]) {
return {
route: {
ipv4_addr: '10.0.0.2',
hostname: 'peer',
version: 'test',
},
peer: {
conns,
},
} as any
}
function peerRoutePairWithDefaultConn(conns: any[], defaultConnId: string) {
const [part1, part2, part3, part4] = defaultConnId
.replaceAll('-', '')
.match(/.{8}/g)!
.map((part) => Number.parseInt(part, 16))
return {
...peerRoutePair(conns),
peer: {
default_conn_id: {
part1,
part2,
part3,
part4,
},
conns,
},
} as any
}
describe('status display helpers', () => {
it('does not render missing IP values as zero addresses', () => {
expect(ipv4ToString(undefined)).toBe('')
expect(ipv4ToString(null)).toBe('')
expect(ipv4ToString({} as any)).toBe('0.0.0.0')
expect(ipv4ToString({ addr: 0 })).toBe('0.0.0.0')
expect(ipv6ToString(undefined)).toBe('')
expect(ipv6ToString(null)).toBe('')
expect(ipv6ToString({} as any)).toBe('::0')
expect(ipv6ToString({ part1: 0, part2: 0, part3: 0, part4: 0 })).toBe('::0')
expect(ipv6ToString({ part4: 1 } as any)).toBe('::1')
})
it('skips missing latency and loss values', () => {
expect(latencyMs(peerRoutePair([
{ conn_id: 'missing', stats: {} },
{ conn_id: 'valid', stats: { latency_us: '2500' } },
{ conn_id: 'invalid', stats: { latency_us: 'unknown' } },
]))).toBe('3ms')
expect(latencyMs(peerRoutePair([
{ conn_id: 'missing', stats: {} },
{ conn_id: 'invalid', stats: { latency_us: 'unknown' } },
]))).toBe('')
expect(lossRate(peerRoutePair([
{ conn_id: 'missing' },
{ conn_id: 'valid', loss_rate: '0.25' },
{ conn_id: 'invalid', loss_rate: 'unknown' },
]))).toBe('25%')
expect(lossRate(peerRoutePair([
{ conn_id: 'missing' },
{ conn_id: 'invalid', loss_rate: 'unknown' },
]))).toBe('')
})
it('prefers the default connection when its metric is valid', () => {
const defaultConnId = '00000001-0002-0003-0004-000000000005'
const conns = [
{ conn_id: 'fallback', stats: { latency_us: '1000' }, loss_rate: '0.01' },
{ conn_id: defaultConnId, stats: { latency_us: '9000' }, loss_rate: '0.5' },
]
expect(latencyMs(peerRoutePairWithDefaultConn(conns, defaultConnId))).toBe('9ms')
expect(lossRate(peerRoutePairWithDefaultConn(conns, defaultConnId))).toBe('50%')
})
})
@@ -1,93 +0,0 @@
import { mount } from '@vue/test-utils'
import { describe, expect, it } from 'vitest'
import { defineComponent, h, nextTick, ref } from 'vue'
import UrlListInput from '../src/components/UrlListInput.vue'
const ButtonStub = defineComponent({
name: 'Button',
emits: ['click'],
setup(_, { slots, emit }) {
return () => h('button', { onClick: (event: MouseEvent) => emit('click', event) }, slots.default?.())
},
})
const UrlInputStub = defineComponent({
name: 'UrlInput',
setup(_, { slots }) {
return () => h('div', slots.actions?.())
},
})
function mountUrlListInput(protos: Record<string, number>, defaultUrl?: string) {
const urls = ref<string[]>([])
const wrapper = mount(defineComponent({
components: { UrlListInput },
setup() {
return { urls, protos, defaultUrl }
},
template: `
<UrlListInput
v-model="urls"
:protos="protos"
:default-url="defaultUrl"
add-label="add_url"
/>
`,
}), {
global: {
stubs: {
Button: ButtonStub,
UrlInput: UrlInputStub,
},
},
})
return { wrapper, urls }
}
describe('UrlListInput.vue add fallback', () => {
it('derives the fallback URL from protos when defaultUrl is not provided', async () => {
const { wrapper, urls } = mountUrlListInput({ tcp: 11010, udp: 11010 })
await wrapper.find('.cursor-pointer').trigger('click')
await nextTick()
expect(urls.value).toEqual(['tcp://0.0.0.0:11010'])
})
it('falls back to the first available protocol when tcp is not present', async () => {
const { wrapper, urls } = mountUrlListInput({ udp: 22000 })
await wrapper.find('.cursor-pointer').trigger('click')
await nextTick()
expect(urls.value).toEqual(['udp://0.0.0.0:22000'])
})
it('falls back to tcp default port when protos is empty', async () => {
const { wrapper, urls } = mountUrlListInput({})
await wrapper.find('.cursor-pointer').trigger('click')
await nextTick()
expect(urls.value).toEqual(['tcp://0.0.0.0:11010'])
})
it('supports port-zero fallback from protos', async () => {
const { wrapper, urls } = mountUrlListInput({ tcp: 0, udp: 0 })
await wrapper.find('.cursor-pointer').trigger('click')
await nextTick()
expect(urls.value).toEqual(['tcp://0.0.0.0:0'])
})
it('uses defaultUrl when provided', async () => {
const { wrapper, urls } = mountUrlListInput({ tcp: 11010 }, 'udp://0.0.0.0:22000')
await wrapper.find('.cursor-pointer').trigger('click')
await nextTick()
expect(urls.value).toEqual(['udp://0.0.0.0:22000'])
})
})
@@ -1,12 +0,0 @@
import { defineConfig } from 'vitest/config'
import vue from '@vitejs/plugin-vue'
import ViteYaml from '@modyfi/vite-plugin-yaml'
export default defineConfig({
plugins: [vue(), ViteYaml()],
test: {
environment: 'happy-dom',
include: ['tests/**/*.spec.ts'],
setupFiles: ['./tests/setup.ts'],
},
})
+3 -3
View File
@@ -4,8 +4,8 @@
"version": "0.0.0",
"type": "module",
"scripts": {
"dev": "pnpm --dir ../frontend-lib build && vite",
"build": "pnpm --dir ../frontend-lib build && vue-tsc -b && vite build",
"dev": "vite",
"build": "vue-tsc -b && vite build",
"preview": "vite preview"
},
"dependencies": {
@@ -32,4 +32,4 @@
"vite-plugin-singlefile": "^2.0.3",
"vue-tsc": "^2.1.10"
}
}
}
+1 -1
View File
@@ -218,7 +218,7 @@ class WebRemoteClient implements Api.RemoteClient {
}
async get_network_info(inst_id: string): Promise<NetworkTypes.NetworkInstanceRunningInfo | undefined> {
const response = await this.client.get<any, Api.CollectNetworkInfoResponse>('/machines/' + this.machine_id + '/networks/info/' + inst_id);
return response.info?.map?.[inst_id];
return response.info.map[inst_id];
}
async list_network_instance_ids(): Promise<Api.ListNetworkInstanceIdResponse> {
const response = await this.client.get<any, ListNetworkInstanceIdResponse>('/machines/' + this.machine_id + '/networks');
-3
View File
@@ -40,9 +40,6 @@ cli:
geoip_db:
en: "The path to the GeoIP2 database file, used to lookup the location of the client, default is the embedded file (only country information) , recommend https://github.com/P3TERX/GeoLite.mmdb"
zh-CN: "GeoIP2 数据库文件路径,用于查找客户端的位置,默认为嵌入文件(仅国家信息),推荐 https://github.com/P3TERX/GeoLite.mmdb"
heartbeat_min_response_ms:
en: "Minimum response time for config-server heartbeat RPCs in milliseconds, default is 0"
zh-CN: "配置服务心跳 RPC 的最短响应时间,单位毫秒,默认为 0"
disable_registration:
en: "Disable user registration"
zh-CN: "禁用用户注册"
File diff suppressed because it is too large Load Diff
+5 -687
View File
@@ -1,5 +1,3 @@
mod managed_config;
mod runtime_reconcile;
pub mod session;
pub mod storage;
@@ -7,7 +5,6 @@ use std::sync::{
Arc,
atomic::{AtomicU32, Ordering},
};
use std::time::Duration;
use dashmap::DashMap;
use easytier::{
@@ -33,10 +30,6 @@ use crate::db::{Db, UserIdInDb, entity::user_running_network_configs};
#[include = "geoip2-cn.mmdb"]
struct GeoipDb;
pub fn is_managed_config_revision_conflict(error: &anyhow::Error) -> bool {
managed_config::is_revision_conflict(error)
}
fn load_geoip_db(geoip_db: Option<String>) -> Option<maxminddb::Reader<Vec<u8>>> {
if let Some(path) = geoip_db {
match maxminddb::Reader::open_readfile(&path) {
@@ -70,14 +63,12 @@ pub struct ClientManager {
webhook_config: SharedWebhookConfig,
geoip_db: Arc<Option<maxminddb::Reader<Vec<u8>>>>,
heartbeat_min_response_delay: Duration,
}
impl ClientManager {
pub fn new(
db: Db,
geoip_db: Option<String>,
heartbeat_min_response_delay: Duration,
feature_flags: Arc<FeatureFlags>,
webhook_config: SharedWebhookConfig,
) -> Self {
@@ -101,7 +92,6 @@ impl ClientManager {
webhook_config,
geoip_db: Arc::new(load_geoip_db(geoip_db)),
heartbeat_min_response_delay,
}
}
@@ -115,7 +105,6 @@ impl ClientManager {
let storage = self.storage.weak_ref();
let listeners_cnt = self.listeners_cnt.clone();
let geoip_db = self.geoip_db.clone();
let heartbeat_min_response_delay = self.heartbeat_min_response_delay;
let feature_flags = self.feature_flags.clone();
let webhook_config = self.webhook_config.clone();
self.tasks.spawn(async move {
@@ -140,7 +129,6 @@ impl ClientManager {
storage.clone(),
client_url.clone(),
location,
heartbeat_min_response_delay,
feature_flags.clone(),
webhook_config.clone(),
);
@@ -161,10 +149,6 @@ impl ClientManager {
self.storage.list_clients()
}
pub async fn list_all_sessions(&self) -> Vec<StorageToken> {
self.storage.list_all_clients()
}
pub fn get_session_by_machine_id(
&self,
user_id: UserIdInDb,
@@ -185,7 +169,7 @@ impl ClientManager {
) -> bool {
let Some(client_url) = self
.storage
.get_client_url_by_machine_id_with_auth(user_id, machine_id, false)
.get_client_url_by_machine_id(user_id, machine_id)
else {
return false;
};
@@ -205,30 +189,14 @@ impl ClientManager {
user_id: UserIdInDb,
machine_id: uuid::Uuid,
desired_configs: Vec<ManagedNetworkConfig>,
config_revision: Option<String>,
expected_config_revision: Option<String>,
) -> anyhow::Result<()> {
let expected_config_revision = match expected_config_revision.as_deref().map(str::trim) {
None => managed_config::ExpectedConfigRevision::Any,
Some("") => managed_config::ExpectedConfigRevision::Exact(None),
Some(revision) => managed_config::ExpectedConfigRevision::Exact(Some(revision)),
};
managed_config::reconcile_web_source_configs(
session::SessionRpcService::reconcile_web_source_configs(
&self.storage,
user_id,
machine_id,
desired_configs,
config_revision.as_deref(),
expected_config_revision,
)
.await?;
if let Some(config_revision) = config_revision
&& let Some(session) = self.get_session_by_machine_id(user_id, &machine_id)
{
session
.notify_config_revision_changed(user_id, machine_id, config_revision)
.await;
}
Ok(())
}
@@ -363,449 +331,19 @@ impl
#[cfg(test)]
mod tests {
use std::{
collections::VecDeque,
sync::{
Arc,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
time::Duration,
};
use std::{sync::Arc, time::Duration};
use axum::{Json, Router, extract::State, routing::post};
use easytier::{
common::MachineIdOptions,
instance_manager::NetworkInstanceManager,
proto::{
api::manage::{NetworkConfig, NetworkingMethod, PortForwardConfig},
common::CompressionAlgoPb,
},
rpc_service::remote_client::Storage as RemoteStorage,
tunnel::{
common::tests::wait_for_condition,
udp::{UdpTunnelConnector, UdpTunnelListener},
},
web_client::{WebClient, run_web_client},
web_client::WebClient,
};
use serde_json::json;
use sqlx::Executor;
use tokio::net::UdpSocket;
use crate::{
FeatureFlags, client_manager::ClientManager, db::Db, webhook::ManagedNetworkConfig,
};
const MANAGED_CONFIG_TOKEN: &str = "managed-config-token";
#[derive(Debug, Clone)]
struct TestWebhookState {
validate_responses: Arc<tokio::sync::Mutex<VecDeque<bool>>>,
validate_count: Arc<AtomicUsize>,
block_second_validate: Arc<AtomicBool>,
allow_second_validate: Arc<AtomicBool>,
}
impl TestWebhookState {
fn new(validate_responses: impl IntoIterator<Item = bool>) -> Self {
Self {
validate_responses: Arc::new(tokio::sync::Mutex::new(
validate_responses.into_iter().collect(),
)),
validate_count: Arc::new(AtomicUsize::new(0)),
block_second_validate: Arc::new(AtomicBool::new(false)),
allow_second_validate: Arc::new(AtomicBool::new(true)),
}
}
fn with_blocked_second_validate(
validate_responses: impl IntoIterator<Item = bool>,
) -> Self {
let state = Self::new(validate_responses);
state.block_second_validate.store(true, Ordering::Release);
state.allow_second_validate.store(false, Ordering::Release);
state
}
fn allow_second_validate(&self) {
self.allow_second_validate.store(true, Ordering::Release);
}
fn validate_count(&self) -> usize {
self.validate_count.load(Ordering::Acquire)
}
}
async fn validate_token_handler(
State(state): State<TestWebhookState>,
) -> Json<serde_json::Value> {
let count = state.validate_count.fetch_add(1, Ordering::AcqRel) + 1;
if count == 2 && state.block_second_validate.load(Ordering::Acquire) {
while !state.allow_second_validate.load(Ordering::Acquire) {
tokio::time::sleep(Duration::from_millis(20)).await;
}
}
let valid = state
.validate_responses
.lock()
.await
.pop_front()
.unwrap_or(true);
if !valid {
return Json(json!({ "valid": false }));
}
Json(json!({
"valid": true,
"binding_version": count,
"config_revision": format!("validated-rev-{count}")
}))
}
async fn webhook_ack_handler() -> Json<serde_json::Value> {
Json(json!({}))
}
async fn test_webhook_config() -> (
crate::webhook::SharedWebhookConfig,
tokio::task::JoinHandle<()>,
TestWebhookState,
) {
let state = TestWebhookState::new([true]);
test_webhook_config_with_state(state).await
}
async fn test_webhook_config_with_state(
state: TestWebhookState,
) -> (
crate::webhook::SharedWebhookConfig,
tokio::task::JoinHandle<()>,
TestWebhookState,
) {
let app = Router::new()
.route("/validate-token", post(validate_token_handler))
.route("/webhook/node-connected", post(webhook_ack_handler))
.route("/webhook/node-disconnected", post(webhook_ack_handler))
.with_state(state.clone());
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
(
Arc::new(crate::webhook::WebhookConfig::new(
Some(format!("http://{addr}")),
None,
None,
None,
None,
)),
server,
state,
)
}
async fn add_random_udp_listener(mgr: &mut ClientManager) -> std::net::SocketAddr {
let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let addr = socket.local_addr().unwrap();
let listener =
UdpTunnelListener::new_with_socket(format!("udp://{addr}").parse().unwrap(), socket);
mgr.add_listener(listener).await.unwrap();
addr
}
async fn wait_for_validated_user(mgr: &ClientManager, machine_id: uuid::Uuid) -> i32 {
tokio::time::timeout(Duration::from_secs(12), async {
loop {
if let Some(token) = mgr.list_sessions().await.into_iter().find(|token| {
token.token == MANAGED_CONFIG_TOKEN && token.machine_id == machine_id
}) {
break token.user_id;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
})
.await
.unwrap()
}
async fn wait_for_validate_count(state: &TestWebhookState, target: usize) {
tokio::time::timeout(Duration::from_secs(12), async {
loop {
if state.validate_count() >= target {
break;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
})
.await
.unwrap();
}
async fn wait_for_session_urls(mgr: &ClientManager) -> Vec<url::Url> {
tokio::time::timeout(Duration::from_secs(12), async {
loop {
let urls = mgr
.client_sessions
.iter()
.map(|entry| entry.key().clone())
.collect::<Vec<_>>();
if !urls.is_empty() {
break urls;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
})
.await
.unwrap()
}
fn managed_config(
instance_id: uuid::Uuid,
network_config: serde_json::Value,
) -> ManagedNetworkConfig {
ManagedNetworkConfig {
instance_id: instance_id.to_string(),
network_config,
}
}
async fn wait_for_runtime_config(
manager: &NetworkInstanceManager,
inst_id: uuid::Uuid,
predicate: impl Fn(&NetworkConfig) -> bool,
) -> NetworkConfig {
tokio::time::timeout(Duration::from_secs(12), async {
loop {
if let Some(config) = manager
.get_instance_config(&inst_id)
.and_then(|config| NetworkConfig::new_from_config(&config).ok())
.filter(|config| predicate(config))
{
break config;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
})
.await
.unwrap()
}
async fn start_web_client_for_test(
config_server_addr: std::net::SocketAddr,
machine_id: uuid::Uuid,
manager: Arc<NetworkInstanceManager>,
) -> WebClient {
run_web_client(
&format!("udp://{config_server_addr}/{MANAGED_CONFIG_TOKEN}"),
MachineIdOptions {
explicit_machine_id: Some(machine_id.to_string()),
state_dir: None,
},
Some("managed-config-core".to_string()),
false,
manager,
None,
)
.await
.unwrap()
}
async fn clear_managed_config_db(
mgr: &ClientManager,
user_id: i32,
machine_id: uuid::Uuid,
instance_id: uuid::Uuid,
) {
mgr.db()
.delete_web_network_configs((user_id, machine_id), &[instance_id])
.await
.unwrap();
sqlx::query("DELETE FROM managed_config_revisions WHERE user_id = ? AND device_id = ?")
.bind(user_id)
.bind(machine_id.to_string())
.execute(&mgr.db().inner())
.await
.unwrap();
}
fn assert_updated_runtime_config(updated: &NetworkConfig, instance_id: uuid::Uuid) {
assert_eq!(
updated.instance_id.as_deref(),
Some(instance_id.to_string().as_str())
);
assert_eq!(updated.dhcp, Some(false));
assert_eq!(updated.virtual_ipv4.as_deref(), Some("10.88.0.7"));
assert_eq!(updated.network_length, Some(24));
assert_eq!(updated.hostname.as_deref(), Some("managed-updated-host"));
assert_eq!(updated.network_name.as_deref(), Some("managed-updated"));
assert_eq!(updated.network_secret.as_deref(), Some("secret-updated"));
assert_eq!(
updated.networking_method,
Some(NetworkingMethod::Manual as i32)
);
assert_eq!(updated.peer_urls, vec!["tcp://127.0.0.1:11010".to_string()]);
assert_eq!(
updated.proxy_cidrs,
vec![
"10.44.0.0/24".to_string(),
"10.45.0.0/24->10.46.0.0/24".to_string()
]
);
assert_eq!(updated.no_tun, Some(true));
assert_eq!(updated.disable_ipv6, Some(true));
assert_eq!(updated.enable_kcp_proxy, Some(true));
assert_eq!(updated.disable_kcp_input, Some(true));
assert_eq!(updated.enable_quic_proxy, Some(true));
assert_eq!(updated.disable_quic_input, Some(true));
assert_eq!(updated.disable_p2p, Some(true));
assert_eq!(updated.p2p_only, Some(true));
assert_eq!(updated.lazy_p2p, Some(true));
assert_eq!(updated.relay_all_peer_rpc, Some(true));
assert_eq!(updated.need_p2p, Some(true));
assert_eq!(updated.multi_thread, Some(false));
assert_eq!(updated.proxy_forward_by_system, Some(true));
assert_eq!(updated.disable_encryption, Some(true));
assert_eq!(updated.enable_relay_network_whitelist, Some(true));
assert_eq!(
updated.relay_network_whitelist,
vec!["10.44.0.0/24".to_string(), "10.45.0.0/24".to_string()]
);
assert_eq!(updated.enable_manual_routes, Some(true));
assert_eq!(
updated.routes,
vec!["10.60.0.0/16".to_string(), "10.61.0.0/16".to_string()]
);
assert_eq!(updated.port_forwards[0].bind_ip, "127.0.0.1");
assert_eq!(updated.port_forwards[0].bind_port, 0);
assert_eq!(updated.port_forwards[0].dst_ip, "10.88.0.8");
assert_eq!(updated.port_forwards[0].dst_port, 80);
assert_eq!(updated.port_forwards[0].proto, "tcp");
assert_eq!(updated.disable_udp_hole_punching, Some(true));
assert_eq!(updated.disable_tcp_hole_punching, Some(true));
assert_eq!(updated.disable_sym_hole_punching, Some(true));
assert_eq!(updated.disable_upnp, Some(true));
assert_eq!(updated.disable_relay_data, Some(true));
assert_eq!(updated.enable_magic_dns, Some(true));
assert_eq!(updated.enable_private_mode, Some(true));
assert_eq!(updated.mtu, Some(1360));
assert_eq!(
updated.data_compress_algo,
Some(CompressionAlgoPb::Zstd as i32)
);
assert_eq!(updated.encryption_algorithm.as_deref(), Some("xor"));
assert_eq!(updated.instance_recv_bps_limit, Some(123456));
assert_eq!(updated.enable_udp_broadcast_relay, Some(true));
assert_eq!(updated.socket_mark, Some(0));
}
fn initial_managed_network_config(inst_id: uuid::Uuid) -> serde_json::Value {
json!({
"instance_id": inst_id.to_string(),
"dhcp": true,
"network_name": "managed-initial",
"network_secret": "secret-initial",
"networking_method": "Standalone",
"no_tun": true,
"disable_ipv6": true,
"enable_kcp_proxy": false,
"disable_kcp_input": false,
"relay_all_peer_rpc": false,
"multi_thread": false,
"disable_relay_data": false,
"mtu": 1380
})
}
fn updated_managed_network_config(inst_id: uuid::Uuid) -> serde_json::Value {
serde_json::to_value(NetworkConfig {
instance_id: Some(inst_id.to_string()),
dhcp: Some(false),
virtual_ipv4: Some("10.88.0.7".to_string()),
network_length: Some(24),
hostname: Some("managed-updated-host".to_string()),
network_name: Some("managed-updated".to_string()),
network_secret: Some("secret-updated".to_string()),
networking_method: Some(NetworkingMethod::Manual as i32),
peer_urls: vec!["tcp://127.0.0.1:11010".to_string()],
proxy_cidrs: vec![
"10.44.0.0/24".to_string(),
"10.45.0.0/24->10.46.0.0/24".to_string(),
],
no_tun: Some(true),
disable_ipv6: Some(true),
enable_kcp_proxy: Some(true),
disable_kcp_input: Some(true),
enable_quic_proxy: Some(true),
disable_quic_input: Some(true),
disable_p2p: Some(true),
p2p_only: Some(true),
lazy_p2p: Some(true),
relay_all_peer_rpc: Some(true),
need_p2p: Some(true),
multi_thread: Some(false),
proxy_forward_by_system: Some(true),
disable_encryption: Some(true),
enable_relay_network_whitelist: Some(true),
relay_network_whitelist: vec!["10.44.0.0/24".to_string(), "10.45.0.0/24".to_string()],
enable_manual_routes: Some(true),
routes: vec!["10.60.0.0/16".to_string(), "10.61.0.0/16".to_string()],
port_forwards: vec![PortForwardConfig {
bind_ip: "127.0.0.1".to_string(),
bind_port: 0,
dst_ip: "10.88.0.8".to_string(),
dst_port: 80,
proto: "tcp".to_string(),
}],
disable_udp_hole_punching: Some(true),
disable_tcp_hole_punching: Some(true),
disable_sym_hole_punching: Some(true),
disable_upnp: Some(true),
disable_relay_data: Some(true),
enable_magic_dns: Some(true),
enable_private_mode: Some(true),
mtu: Some(1360),
data_compress_algo: Some(CompressionAlgoPb::Zstd as i32),
encryption_algorithm: Some("xor".to_string()),
instance_recv_bps_limit: Some(123456),
enable_udp_broadcast_relay: Some(true),
socket_mark: Some(0),
..Default::default()
})
.unwrap()
}
fn redelivered_managed_network_config(inst_id: uuid::Uuid) -> serde_json::Value {
serde_json::to_value(NetworkConfig {
instance_id: Some(inst_id.to_string()),
dhcp: Some(false),
virtual_ipv4: Some("10.88.0.7".to_string()),
network_length: Some(24),
hostname: Some("managed-redelivered-host".to_string()),
network_name: Some("managed-redelivered".to_string()),
network_secret: Some("secret-updated".to_string()),
networking_method: Some(NetworkingMethod::Manual as i32),
peer_urls: vec!["tcp://127.0.0.1:11010".to_string()],
proxy_cidrs: vec![
"10.44.0.0/24".to_string(),
"10.45.0.0/24->10.46.0.0/24".to_string(),
],
no_tun: Some(true),
disable_ipv6: Some(true),
enable_kcp_proxy: Some(true),
disable_kcp_input: Some(true),
relay_all_peer_rpc: Some(true),
need_p2p: Some(true),
multi_thread: Some(false),
enable_private_mode: Some(true),
mtu: Some(1360),
data_compress_algo: Some(CompressionAlgoPb::Zstd as i32),
encryption_algorithm: Some("xor".to_string()),
instance_recv_bps_limit: Some(654321),
..Default::default()
})
.unwrap()
}
use crate::{FeatureFlags, client_manager::ClientManager, db::Db};
#[tokio::test]
async fn test_client() {
@@ -813,7 +351,6 @@ mod tests {
let mut mgr = ClientManager::new(
Db::memory_db().await,
None,
Duration::ZERO,
Arc::new(FeatureFlags::default()),
Arc::new(crate::webhook::WebhookConfig::new(
None, None, None, None, None,
@@ -873,223 +410,4 @@ mod tests {
println!("{:?}", req);
println!("{:?}", mgr);
}
#[tokio::test]
async fn managed_web_config_revision_updates_running_core_config() {
let (webhook_config, webhook_server, _) = test_webhook_config().await;
let mut mgr = ClientManager::new(
Db::memory_db().await,
None,
Duration::ZERO,
Arc::new(FeatureFlags::default()),
webhook_config,
);
let config_server_addr = add_random_udp_listener(&mut mgr).await;
let machine_id = uuid::Uuid::new_v4();
let instance_id = uuid::Uuid::new_v4();
let core_manager = Arc::new(NetworkInstanceManager::new());
let client =
start_web_client_for_test(config_server_addr, machine_id, core_manager.clone()).await;
let user_id = wait_for_validated_user(&mgr, machine_id).await;
mgr.reconcile_managed_network_configs(
user_id,
machine_id,
vec![managed_config(
instance_id,
initial_managed_network_config(instance_id),
)],
Some("rev-initial".to_string()),
None,
)
.await
.unwrap();
wait_for_runtime_config(&core_manager, instance_id, |config| {
config.network_name.as_deref() == Some("managed-initial")
})
.await;
// Online revision update: web-owned running config is fully overwritten
// when non-hot-patch flags such as enable_kcp_proxy change.
mgr.reconcile_managed_network_configs(
user_id,
machine_id,
vec![managed_config(
instance_id,
updated_managed_network_config(instance_id),
)],
Some("rev-updated".to_string()),
Some("rev-initial".to_string()),
)
.await
.unwrap();
let updated = wait_for_runtime_config(&core_manager, instance_id, |config| {
config.network_name.as_deref() == Some("managed-updated")
&& config.enable_kcp_proxy == Some(true)
&& config.port_forwards.len() == 1
})
.await;
assert_updated_runtime_config(&updated, instance_id);
assert_eq!(
core_manager.get_instance_network_config_source(&instance_id),
Some(easytier::common::config::ConfigSource::Web)
);
assert_eq!(
mgr.db()
.get_managed_config_revision((user_id, machine_id))
.await
.unwrap()
.as_deref(),
Some("rev-updated")
);
// Web DB loss path: clear web-owned config and revision, then simulate
// the webhook re-posting the authoritative desired config. The already
// connected session should receive the distinguishable re-delivered
// revision without restarting.
clear_managed_config_db(&mgr, user_id, machine_id, instance_id).await;
assert!(
mgr.db()
.get_network_config((user_id, machine_id), &instance_id.to_string())
.await
.unwrap()
.is_none()
);
assert!(
mgr.db()
.get_managed_config_revision((user_id, machine_id))
.await
.unwrap()
.is_none()
);
mgr.reconcile_managed_network_configs(
user_id,
machine_id,
vec![managed_config(
instance_id,
redelivered_managed_network_config(instance_id),
)],
Some("rev-webhook-redelivery".to_string()),
None,
)
.await
.unwrap();
let redelivered = wait_for_runtime_config(&core_manager, instance_id, |config| {
config.network_name.as_deref() == Some("managed-redelivered")
&& config.instance_recv_bps_limit == Some(654321)
})
.await;
assert_eq!(
redelivered.instance_id.as_deref(),
Some(instance_id.to_string().as_str())
);
assert_eq!(
redelivered.hostname.as_deref(),
Some("managed-redelivered-host")
);
assert_eq!(
redelivered.network_name.as_deref(),
Some("managed-redelivered")
);
assert_eq!(redelivered.enable_kcp_proxy, Some(true));
assert_eq!(redelivered.instance_recv_bps_limit, Some(654321));
assert_eq!(
core_manager.get_instance_network_config_source(&instance_id),
Some(easytier::common::config::ConfigSource::Web)
);
assert_eq!(
mgr.db()
.get_managed_config_revision((user_id, machine_id))
.await
.unwrap()
.as_deref(),
Some("rev-webhook-redelivery")
);
// Reconnect path: a fresh core manager has no local runtime state, so
// the new session must replay the managed config persisted in web DB.
drop(client);
let reconnected_core_manager = Arc::new(NetworkInstanceManager::new());
let _reconnected_client = start_web_client_for_test(
config_server_addr,
machine_id,
reconnected_core_manager.clone(),
)
.await;
wait_for_validated_user(&mgr, machine_id).await;
let replayed = wait_for_runtime_config(&reconnected_core_manager, instance_id, |config| {
config.network_name.as_deref() == Some("managed-redelivered")
&& config.instance_recv_bps_limit == Some(654321)
})
.await;
assert_eq!(
replayed.network_name.as_deref(),
Some("managed-redelivered")
);
assert_eq!(replayed.enable_kcp_proxy, Some(true));
assert_eq!(replayed.instance_recv_bps_limit, Some(654321));
webhook_server.abort();
}
#[tokio::test]
async fn webhook_reject_disconnects_and_revalidates_after_reconnect() {
let webhook_state = TestWebhookState::with_blocked_second_validate([false, true]);
let (webhook_config, webhook_server, webhook_state) =
test_webhook_config_with_state(webhook_state).await;
let mut mgr = ClientManager::new(
Db::memory_db().await,
None,
Duration::ZERO,
Arc::new(FeatureFlags::default()),
webhook_config,
);
let config_server_addr = add_random_udp_listener(&mut mgr).await;
let machine_id = uuid::Uuid::new_v4();
let core_manager = Arc::new(NetworkInstanceManager::new());
let client =
start_web_client_for_test(config_server_addr, machine_id, core_manager.clone()).await;
let first_session_urls = wait_for_session_urls(&mgr).await;
wait_for_validate_count(&webhook_state, 1).await;
wait_for_validate_count(&webhook_state, 2).await;
assert!(
mgr.list_sessions().await.is_empty(),
"invalid validate-token response must not authorize the session"
);
webhook_state.allow_second_validate();
let user_id = wait_for_validated_user(&mgr, machine_id).await;
tokio::time::timeout(Duration::from_secs(12), async {
loop {
let reconnected = mgr
.client_sessions
.iter()
.any(|entry| !first_session_urls.iter().any(|url| url == entry.key()));
if reconnected {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
})
.await
.unwrap();
assert!(
client.is_connected(),
"web client should reconnect after invalid session heartbeat failure"
);
assert!(webhook_state.validate_count() >= 2);
assert!(
mgr.get_session_by_machine_id(user_id, &machine_id)
.is_some()
);
webhook_server.abort();
}
}
@@ -1,812 +0,0 @@
use anyhow::Context as _;
use easytier::{
common::config::{
ConfigLoader, EncryptionAlgorithm, PortForwardConfig as RuntimePortForwardConfig,
},
proto::{
acl::Acl,
api::{
config::{
AclPatch, ConfigPatchAction, InstanceConfigPatch, PatchConfigRequest,
PortForwardPatch, ProxyNetworkPatch,
},
instance::{InstanceIdentifier, instance_identifier},
manage::{
ConfigSource as RpcConfigSource, GetNetworkInstanceConfigRequest, NetworkConfig,
RunNetworkInstanceRequest,
},
},
common::{CompressionAlgoPb, Ipv4Inet as RpcIpv4Inet},
rpc_types::controller::BaseController,
},
};
use super::session::{SessionConfigClient, SessionRpcClient};
pub(super) enum RuntimeReconcileAction {
None,
Run {
config: Box<NetworkConfig>,
overwrite: bool,
},
Patch(Box<InstanceConfigPatch>),
}
#[derive(Clone, PartialEq)]
struct RuntimeProxyNetwork {
cidr: String,
mapped_cidr: Option<String>,
}
fn instance_identifier(inst_id: &str) -> anyhow::Result<InstanceIdentifier> {
let inst_id = uuid::Uuid::parse_str(inst_id)
.with_context(|| format!("invalid runtime instance id: {inst_id}"))?;
Ok(InstanceIdentifier {
selector: Some(instance_identifier::Selector::Id(inst_id.into())),
})
}
fn hot_patch_base(config: &NetworkConfig) -> anyhow::Result<NetworkConfig> {
let data_compress_algo = normalized_data_compress_algo(config.data_compress_algo);
let encryption_algorithm = normalized_encryption_algorithm(config.encryption_algorithm.clone());
let mut config = NetworkConfig::new_from_config(config.gen_config()?)?;
let is_credential_mode = config.network_secret.is_none()
&& config
.secure_mode
.as_ref()
.and_then(|mode| mode.local_private_key.as_deref())
.is_some_and(|key| !key.is_empty());
config.acl = None;
config.port_forwards.clear();
config.proxy_cidrs.clear();
config.disable_relay_data = None;
if config.dhcp.unwrap_or_default() {
config.virtual_ipv4 = None;
config.network_length = None;
}
if let Some(secure_mode) = config.secure_mode.as_mut() {
if !is_credential_mode {
secure_mode.local_private_key = None;
}
secure_mode.local_public_key = None;
}
config.data_compress_algo = data_compress_algo;
config.encryption_algorithm = encryption_algorithm;
Ok(config)
}
fn normalized_data_compress_algo(algo: Option<i32>) -> Option<i32> {
let default = CompressionAlgoPb::None as i32;
let effective = algo.map(|algo| if algo < default { default } else { algo });
effective.filter(|algo| *algo != default)
}
fn normalized_encryption_algorithm(algo: Option<String>) -> Option<String> {
let default = EncryptionAlgorithm::default().to_string();
algo.filter(|algo| algo != &default)
}
fn diff_port_forwards(
current: &[RuntimePortForwardConfig],
desired: &[RuntimePortForwardConfig],
) -> Vec<PortForwardPatch> {
let mut patches = Vec::new();
for cfg in unique_port_forwards(current, desired) {
let current_count = current.iter().filter(|item| *item == &cfg).count();
let desired_count = desired.iter().filter(|item| *item == &cfg).count();
if current_count == desired_count {
continue;
}
if current_count > 0 {
patches.push(PortForwardPatch {
action: ConfigPatchAction::Remove as i32,
cfg: Some(cfg.clone().into()),
});
}
patches.extend((0..desired_count).map(|_| PortForwardPatch {
action: ConfigPatchAction::Add as i32,
cfg: Some(cfg.clone().into()),
}));
}
patches
}
fn unique_port_forwards(
current: &[RuntimePortForwardConfig],
desired: &[RuntimePortForwardConfig],
) -> Vec<RuntimePortForwardConfig> {
let mut unique = Vec::new();
for cfg in current.iter().chain(desired.iter()) {
if !unique.contains(cfg) {
unique.push(cfg.clone());
}
}
unique
}
fn parse_rpc_ipv4_inet(value: &str) -> anyhow::Result<RpcIpv4Inet> {
value
.parse::<RpcIpv4Inet>()
.with_context(|| format!("failed to parse runtime ipv4 cidr: {value}"))
}
fn diff_proxy_networks(
current: &[RuntimeProxyNetwork],
desired: &[RuntimeProxyNetwork],
) -> anyhow::Result<Vec<ProxyNetworkPatch>> {
if current == desired {
return Ok(Vec::new());
}
let mut patches = vec![ProxyNetworkPatch {
action: ConfigPatchAction::Clear as i32,
cidr: Some(clear_proxy_network_cidr(current, desired)?),
..Default::default()
}];
for proxy_network in desired {
patches.push(ProxyNetworkPatch {
action: ConfigPatchAction::Add as i32,
cidr: Some(parse_rpc_ipv4_inet(&proxy_network.cidr)?),
mapped_cidr: proxy_network
.mapped_cidr
.as_deref()
.map(parse_rpc_ipv4_inet)
.transpose()?,
});
}
Ok(patches)
}
fn clear_proxy_network_cidr(
current: &[RuntimeProxyNetwork],
desired: &[RuntimeProxyNetwork],
) -> anyhow::Result<RpcIpv4Inet> {
let cidr = desired
.first()
.or_else(|| current.first())
.map(|proxy_network| proxy_network.cidr.as_str())
.unwrap_or("0.0.0.0/0");
parse_rpc_ipv4_inet(cidr)
}
fn normalized_acl(acl: &Option<Acl>) -> Option<Acl> {
let acl = acl.clone().unwrap_or_default();
(acl != Acl::default()).then_some(acl)
}
fn normalized_port_forwards(
config: &NetworkConfig,
) -> anyhow::Result<Vec<RuntimePortForwardConfig>> {
Ok(config
.gen_config()?
.get_port_forwards()
.into_iter()
.map(|cfg| {
RuntimePortForwardConfig::from(easytier::proto::common::PortForwardConfigPb::from(cfg))
})
.collect())
}
fn normalized_proxy_networks(config: &NetworkConfig) -> anyhow::Result<Vec<RuntimeProxyNetwork>> {
Ok(config
.gen_config()?
.get_proxy_cidrs()
.into_iter()
.map(|proxy_network| RuntimeProxyNetwork {
cidr: proxy_network.cidr.to_string(),
mapped_cidr: proxy_network.mapped_cidr.map(|cidr| cidr.to_string()),
})
.collect())
}
fn normalized_disable_relay_data(config: &NetworkConfig) -> anyhow::Result<bool> {
Ok(config.gen_config()?.get_flags().disable_relay_data)
}
fn web_source_runtime_patch(
current: &NetworkConfig,
desired: &NetworkConfig,
) -> anyhow::Result<Option<InstanceConfigPatch>> {
if let Some(desired_hostname) = desired
.hostname
.as_deref()
.filter(|hostname| !hostname.is_empty())
&& current.hostname.as_deref() != Some(desired_hostname)
{
return Ok(None);
}
let mut current_base = hot_patch_base(current)?;
let mut desired_base = hot_patch_base(desired)?;
current_base.hostname = None;
desired_base.hostname = None;
if current_base != desired_base {
return Ok(None);
}
let mut patch = InstanceConfigPatch::default();
let current_acl = normalized_acl(&current.acl);
let desired_acl = normalized_acl(&desired.acl);
if current_acl != desired_acl {
patch.acl = Some(AclPatch {
acl: Some(desired_acl.unwrap_or_default()),
..Default::default()
});
}
let current_port_forwards = normalized_port_forwards(current)?;
let desired_port_forwards = normalized_port_forwards(desired)?;
if current_port_forwards != desired_port_forwards {
patch.port_forwards = diff_port_forwards(&current_port_forwards, &desired_port_forwards);
}
let current_proxy_networks = normalized_proxy_networks(current)?;
let desired_proxy_networks = normalized_proxy_networks(desired)?;
if current_proxy_networks != desired_proxy_networks {
if current_proxy_networks.is_empty() {
return Ok(None);
}
patch.proxy_networks =
diff_proxy_networks(&current_proxy_networks, &desired_proxy_networks)?;
}
let current_disable_relay_data = normalized_disable_relay_data(current)?;
let desired_disable_relay_data = normalized_disable_relay_data(desired)?;
if current_disable_relay_data != desired_disable_relay_data {
patch.disable_relay_data = Some(desired_disable_relay_data);
}
Ok(Some(patch))
}
fn ensure_runtime_config_converged(
current: &NetworkConfig,
desired: &NetworkConfig,
) -> anyhow::Result<()> {
let patch = web_source_runtime_patch(current, desired)?;
match patch {
Some(patch) if patch == InstanceConfigPatch::default() => Ok(()),
Some(patch) => anyhow::bail!("runtime config still needs patch after reconcile: {patch:?}"),
None => anyhow::bail!("runtime config still needs full overwrite after reconcile"),
}
}
async fn run_web_source_instance(
rpc_client: &mut SessionRpcClient,
inst_id: &str,
config: NetworkConfig,
overwrite: bool,
) -> anyhow::Result<()> {
rpc_client
.run_network_instance(
BaseController::default(),
RunNetworkInstanceRequest {
inst_id: Some(inst_id.to_string().into()),
config: Some(config),
overwrite,
source: RpcConfigSource::Web as i32,
},
)
.await?;
Ok(())
}
pub(super) async fn get_runtime_config(
rpc_client: &mut SessionRpcClient,
inst_id: &str,
) -> anyhow::Result<NetworkConfig> {
rpc_client
.get_network_instance_config(
BaseController::default(),
GetNetworkInstanceConfigRequest {
inst_id: Some(inst_id.to_string().into()),
},
)
.await?
.config
.ok_or_else(|| anyhow::anyhow!("runtime returned empty config for {inst_id}"))
}
pub(super) async fn prepare_web_source_runtime_reconcile(
rpc_client: &mut SessionRpcClient,
inst_id: &str,
desired_config: NetworkConfig,
is_running: bool,
) -> anyhow::Result<RuntimeReconcileAction> {
if !is_running {
return Ok(RuntimeReconcileAction::Run {
config: Box::new(desired_config),
overwrite: false,
});
}
let current_config = get_runtime_config(rpc_client, inst_id).await?;
prepare_web_source_runtime_reconcile_from_current(&current_config, desired_config)
}
pub(super) fn prepare_web_source_runtime_reconcile_from_current(
current_config: &NetworkConfig,
desired_config: NetworkConfig,
) -> anyhow::Result<RuntimeReconcileAction> {
let Some(patch) = web_source_runtime_patch(current_config, &desired_config)? else {
return Ok(RuntimeReconcileAction::Run {
config: Box::new(desired_config),
overwrite: true,
});
};
if patch == InstanceConfigPatch::default() {
return Ok(RuntimeReconcileAction::None);
}
Ok(RuntimeReconcileAction::Patch(Box::new(patch)))
}
pub(super) async fn apply_web_source_runtime_reconcile(
rpc_client: &mut SessionRpcClient,
config_client: &mut SessionConfigClient,
inst_id: &str,
desired_config: NetworkConfig,
action: RuntimeReconcileAction,
) -> anyhow::Result<NetworkConfig> {
match action {
RuntimeReconcileAction::None => Ok(desired_config),
RuntimeReconcileAction::Run { config, overwrite } => {
run_web_source_instance(rpc_client, inst_id, *config, overwrite).await?;
Ok(desired_config)
}
RuntimeReconcileAction::Patch(patch) => {
config_client
.patch_config(
BaseController::default(),
PatchConfigRequest {
instance: Some(instance_identifier(inst_id)?),
patch: Some(*patch),
},
)
.await?;
let current_config = get_runtime_config(rpc_client, inst_id).await?;
ensure_runtime_config_converged(&current_config, &desired_config)?;
Ok(current_config)
}
}
}
#[cfg(test)]
mod tests {
use easytier::proto::{
api::{
config::ConfigPatchAction,
manage::{NetworkingMethod, PortForwardConfig},
},
common::{CompressionAlgoPb, SocketType},
};
use super::*;
fn config_with_port_forwards(port_forwards: Vec<PortForwardConfig>) -> NetworkConfig {
NetworkConfig {
instance_id: Some("11111111-1111-1111-1111-111111111111".to_string()),
dhcp: Some(true),
network_name: Some("managed".to_string()),
network_secret: Some("secret".to_string()),
networking_method: Some(NetworkingMethod::Manual as i32),
port_forwards,
..Default::default()
}
}
fn port_forward(bind_port: u32, dst_port: u32) -> PortForwardConfig {
PortForwardConfig {
bind_ip: "127.0.0.1".to_string(),
bind_port,
dst_ip: "10.144.0.1".to_string(),
dst_port,
proto: "tcp".to_string(),
}
}
fn patch_port(patch: &PortForwardPatch) -> (i32, u32, u32, i32) {
let cfg = patch.cfg.as_ref().expect("port forward patch cfg");
(
patch.action,
cfg.bind_addr.as_ref().expect("bind addr").port,
cfg.dst_addr.as_ref().expect("dst addr").port,
cfg.socket_type,
)
}
fn patch_proxy_network(patch: &ProxyNetworkPatch) -> (i32, String, Option<String>) {
(
patch.action,
patch.cidr.map(|cidr| cidr.to_string()).unwrap_or_default(),
patch.mapped_cidr.map(|cidr| cidr.to_string()),
)
}
#[test]
fn runtime_patch_ignores_runtime_defaults_and_adds_port_forward() {
let mut current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
current.virtual_ipv4 = Some("10.144.0.2".to_string());
current.network_length = Some(16);
current.bind_device = Some(true);
current.dev_name = Some(String::new());
current.disable_ipv6 = Some(false);
current.mtu = Some(1380);
current.multi_thread = Some(true);
let desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.port_forwards.len(), 1);
assert_eq!(
patch_port(&patch.port_forwards[0]),
(
ConfigPatchAction::Add as i32,
23007,
3389,
SocketType::Tcp as i32
)
);
}
#[test]
fn runtime_patch_removes_deleted_port_forward_without_clear() {
let current =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
let desired = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.port_forwards.len(), 1);
assert_eq!(
patch_port(&patch.port_forwards[0]),
(
ConfigPatchAction::Remove as i32,
23007,
3389,
SocketType::Tcp as i32
)
);
}
#[test]
fn runtime_patch_reconciles_duplicate_port_forward_count() {
let current =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23000, 5174)]);
let desired = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.port_forwards.len(), 2);
assert_eq!(
patch_port(&patch.port_forwards[0]),
(
ConfigPatchAction::Remove as i32,
23000,
5174,
SocketType::Tcp as i32
)
);
assert_eq!(
patch_port(&patch.port_forwards[1]),
(
ConfigPatchAction::Add as i32,
23000,
5174,
SocketType::Tcp as i32
)
);
}
#[test]
fn runtime_convergence_rejects_stale_extra_port_forward() {
let current = config_with_port_forwards(vec![
port_forward(23000, 5174),
port_forward(23007, 3389),
port_forward(23100, 8080),
]);
let desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
let err = ensure_runtime_config_converged(&current, &desired)
.expect_err("extra runtime port forward should not converge");
assert!(
err.to_string()
.contains("runtime config still needs patch after reconcile"),
"unexpected error: {err:?}"
);
}
#[test]
fn runtime_patch_canonicalizes_port_forward_protocol() {
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let mut desired_port_forward = port_forward(23000, 5174);
desired_port_forward.proto = "TCP".to_string();
let desired = config_with_port_forwards(vec![desired_port_forward]);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch, InstanceConfigPatch::default());
ensure_runtime_config_converged(&current, &desired).expect("runtime converged");
}
#[test]
fn runtime_patch_rejects_non_hot_config_change() {
let current = config_with_port_forwards(Vec::new());
let mut desired = current.clone();
desired.network_secret = Some("new-secret".to_string());
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_rejects_routes_change() {
let mut current = config_with_port_forwards(Vec::new());
current.enable_manual_routes = Some(true);
current.routes = vec!["10.1.0.0/16".to_string(), "10.2.0.0/16".to_string()];
let mut desired = config_with_port_forwards(Vec::new());
desired.enable_manual_routes = Some(true);
desired.routes = vec!["10.2.0.0/16".to_string(), "10.3.0.0/16".to_string()];
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_replaces_proxy_networks() {
let mut current = config_with_port_forwards(Vec::new());
current.proxy_cidrs = vec![
"10.1.0.0/16".to_string(),
"10.2.0.0/16->10.20.0.0/16".to_string(),
];
let mut desired = config_with_port_forwards(Vec::new());
desired.proxy_cidrs = vec![
"10.2.0.0/16->10.21.0.0/16".to_string(),
"10.3.0.0/16".to_string(),
];
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.proxy_networks.len(), 3);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[0]),
(
ConfigPatchAction::Clear as i32,
"10.2.0.0/16".to_string(),
None
)
);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[1]),
(
ConfigPatchAction::Add as i32,
"10.2.0.0/16".to_string(),
Some("10.21.0.0/16".to_string())
)
);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[2]),
(
ConfigPatchAction::Add as i32,
"10.3.0.0/16".to_string(),
None
)
);
}
#[test]
fn runtime_patch_replaces_proxy_networks_with_same_source_cidr() {
let mut current = config_with_port_forwards(Vec::new());
current.proxy_cidrs = vec![
"10.1.2.0/24".to_string(),
"10.1.2.0/24->10.1.3.0/24".to_string(),
];
let mut desired = config_with_port_forwards(Vec::new());
desired.proxy_cidrs = vec!["10.1.2.0/24->10.1.3.0/24".to_string()];
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.proxy_networks.len(), 2);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[0]),
(
ConfigPatchAction::Clear as i32,
"10.1.2.0/24".to_string(),
None
)
);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[1]),
(
ConfigPatchAction::Add as i32,
"10.1.2.0/24".to_string(),
Some("10.1.3.0/24".to_string())
)
);
}
#[test]
fn runtime_patch_rejects_proxy_network_empty_to_nonempty() {
let current = config_with_port_forwards(Vec::new());
let mut desired = config_with_port_forwards(Vec::new());
desired.proxy_cidrs = vec!["10.1.2.0/24".to_string()];
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_clears_proxy_networks_with_legacy_compatible_cidr() {
let mut current = config_with_port_forwards(Vec::new());
current.proxy_cidrs = vec!["10.1.2.0/24".to_string()];
let desired = config_with_port_forwards(Vec::new());
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.proxy_networks.len(), 1);
assert_eq!(
patch_proxy_network(&patch.proxy_networks[0]),
(
ConfigPatchAction::Clear as i32,
"10.1.2.0/24".to_string(),
None
)
);
}
#[test]
fn runtime_patch_updates_disable_relay_data() {
let current = config_with_port_forwards(Vec::new());
let mut desired = current.clone();
desired.disable_relay_data = Some(true);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.disable_relay_data, Some(true));
}
#[test]
fn runtime_patch_still_rejects_unsupported_flag_change() {
let current = config_with_port_forwards(Vec::new());
let mut desired = current.clone();
desired.no_tun = Some(true);
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_rejects_encryption_algorithm_change() {
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let mut desired = current.clone();
desired.encryption_algorithm = Some("managed-test-algo".to_string());
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_rejects_data_compress_algo_change() {
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let mut desired = current.clone();
desired.data_compress_algo = Some(CompressionAlgoPb::Zstd as i32);
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_rejects_credential_private_key_change() {
let mut current = config_with_port_forwards(Vec::new());
current.network_secret = None;
current.secure_mode = Some(easytier::proto::common::SecureModeConfig {
enabled: true,
local_private_key: Some("mUuD5fsIm/ftvgS4WBAYFMNLqWX3qT9rnm4PrnOqb9s=".to_string()),
local_public_key: None,
});
let mut desired = current.clone();
desired.secure_mode = Some(easytier::proto::common::SecureModeConfig {
enabled: true,
local_private_key: Some("aEpz80FuYbaY4QLJizAIuIcK4TYsoSA9jHHCXCOQJoc=".to_string()),
local_public_key: None,
});
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
#[test]
fn runtime_patch_ignores_generated_secure_key_when_network_secret_exists() {
let mut current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
current.secure_mode = Some(easytier::proto::common::SecureModeConfig {
enabled: true,
local_private_key: Some("mUuD5fsIm/ftvgS4WBAYFMNLqWX3qT9rnm4PrnOqb9s=".to_string()),
local_public_key: Some("4x6L5dZjB8hsPO4f96Hyhi4xFealBu6i3BxRVBYR1Fc=".to_string()),
});
let mut desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
desired.secure_mode = Some(easytier::proto::common::SecureModeConfig {
enabled: true,
local_private_key: None,
local_public_key: None,
});
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.port_forwards.len(), 1);
assert_eq!(
patch_port(&patch.port_forwards[0]),
(
ConfigPatchAction::Add as i32,
23007,
3389,
SocketType::Tcp as i32
)
);
}
#[test]
fn runtime_patch_ignores_runtime_hostname_when_desired_omits_hostname() {
let mut current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
current.hostname = Some("runtime-host".to_string());
let desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
let patch = web_source_runtime_patch(&current, &desired)
.expect("build patch")
.expect("hot patch");
assert_eq!(patch.port_forwards.len(), 1);
assert_eq!(
patch_port(&patch.port_forwards[0]),
(
ConfigPatchAction::Add as i32,
23007,
3389,
SocketType::Tcp as i32
)
);
}
#[test]
fn runtime_patch_rejects_explicit_desired_hostname_change() {
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let mut desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
desired.hostname =
Some(easytier::common::config::TomlConfigLoader::default().get_hostname());
let patch = web_source_runtime_patch(&current, &desired).expect("build patch");
assert!(patch.is_none());
}
}
File diff suppressed because it is too large Load Diff
@@ -1,913 +0,0 @@
use std::collections::{HashMap, HashSet};
use easytier::{
proto::{
api::manage::{
DeleteNetworkInstanceRequest, ListNetworkInstanceMetaRequest,
ListNetworkInstanceRequest, NetworkConfig, NetworkMeta, RunNetworkInstanceRequest,
},
rpc_types::controller::BaseController,
web::HeartbeatRequest,
},
rpc_service::remote_client::{ListNetworkProps, Storage as _},
};
use tokio::sync::{RwLock, broadcast};
use super::{SessionConfigClient, SessionData, SessionRpcClient, SessionRpcService};
use crate::client_manager::{
managed_config::{self, PersistedConfigSource},
runtime_reconcile,
storage::{StorageInner, WeakRefStorage},
};
async fn recv_latest_heartbeat(
heartbeat_waiter: &mut broadcast::Receiver<HeartbeatRequest>,
) -> Option<HeartbeatRequest> {
let mut req = loop {
match heartbeat_waiter.recv().await {
Ok(req) => break req,
Err(broadcast::error::RecvError::Lagged(skipped)) => {
tracing::warn!(
skipped,
"heartbeat reconcile worker lagged, waiting for latest request"
);
}
Err(broadcast::error::RecvError::Closed) => {
tracing::error!("Failed to receive heartbeat request: channel closed");
return None;
}
}
};
// Drop any heartbeat backlog accumulated while the previous reconcile
// round was doing DB/RPC IO. The newest heartbeat has the freshest
// runtime instance list, which is all this task needs.
loop {
match heartbeat_waiter.try_recv() {
Ok(next_req) => req = next_req,
Err(broadcast::error::TryRecvError::Empty) => break,
Err(broadcast::error::TryRecvError::Lagged(_)) => continue,
Err(broadcast::error::TryRecvError::Closed) => return None,
}
}
Some(req)
}
pub(super) async fn reconcile_network_configs_on_heartbeat(
session_data: std::sync::Weak<RwLock<SessionData>>,
mut heartbeat_waiter: broadcast::Receiver<HeartbeatRequest>,
storage: WeakRefStorage,
mut rpc_client: SessionRpcClient,
mut config_client: SessionConfigClient,
) {
let mut cache = ReconcileCache::default();
loop {
let Some(req) = recv_latest_heartbeat(&mut heartbeat_waiter).await else {
return;
};
let Some(storage) = storage.upgrade() else {
tracing::error!("Failed to get storage");
return;
};
let mut round =
match prepare_reconcile_round(&session_data, &storage, &mut rpc_client, req).await {
RoundStatus::Ready(round) => round,
RoundStatus::Skip => continue,
RoundStatus::Stop => return,
};
let running_metas =
match sync_running_sources_for_round(&mut rpc_client, &storage, &mut round).await {
RoundStatus::Ready(running_metas) => running_metas,
RoundStatus::Skip => continue,
RoundStatus::Stop => return,
};
let desired_web_inst_ids =
managed_config::desired_web_source_instance_ids(&round.local_configs);
cache.runtime_configs.retain_desired(&desired_web_inst_ids);
let mut outcome = match cleanup_stale_web_source_instances(
&session_data,
&storage,
&mut rpc_client,
&round,
running_metas.as_deref(),
&desired_web_inst_ids,
&mut cache,
)
.await
{
RoundStatus::Ready(outcome) => outcome,
RoundStatus::Skip => continue,
RoundStatus::Stop => return,
};
outcome.merge(
reconcile_desired_runtime_configs(
&session_data,
&mut rpc_client,
&mut config_client,
&round,
&mut cache,
)
.await,
);
if !outcome.has_failed {
cache.last_desired_web_inst_ids = Some(desired_web_inst_ids);
}
match mark_config_revision_applied_if_current(&session_data, &storage, &round, &outcome)
.await
{
RoundStatus::Ready(()) | RoundStatus::Skip => {}
RoundStatus::Stop => return,
}
}
}
enum RoundStatus<T> {
Ready(T),
Skip,
Stop,
}
enum ConfigActionResult {
Success,
Failed,
StopRound,
}
#[derive(Default)]
struct ReconcileCache {
cleaned_web_source_instances: bool,
last_desired_web_inst_ids: Option<HashSet<String>>,
runtime_configs: SessionRuntimeConfigCache,
}
#[derive(Default)]
struct SessionRuntimeConfigCache {
entries: HashMap<String, NetworkConfig>,
}
impl SessionRuntimeConfigCache {
fn plan(
&self,
inst_id: &str,
desired_config: NetworkConfig,
) -> anyhow::Result<Option<runtime_reconcile::RuntimeReconcileAction>> {
let Some(observed_config) = self.entries.get(inst_id) else {
return Ok(None);
};
runtime_reconcile::prepare_web_source_runtime_reconcile_from_current(
observed_config,
desired_config,
)
.map(Some)
}
fn remember(&mut self, inst_id: &str, observed_config: NetworkConfig) {
self.entries.insert(inst_id.to_string(), observed_config);
}
fn forget(&mut self, inst_id: &str) {
self.entries.remove(inst_id);
}
fn forget_many<'a>(&mut self, inst_ids: impl IntoIterator<Item = &'a String>) {
for inst_id in inst_ids {
self.entries.remove(inst_id);
}
}
fn retain_desired(&mut self, desired_web_inst_ids: &HashSet<String>) {
self.entries
.retain(|inst_id, _| desired_web_inst_ids.contains(inst_id));
}
}
#[derive(Default)]
struct ReconcileOutcome {
has_failed: bool,
managed_revision_failed: bool,
}
impl ReconcileOutcome {
fn record_failure(&mut self, managed_revision_failed: bool) {
self.has_failed = true;
self.managed_revision_failed |= managed_revision_failed;
}
fn merge(&mut self, other: Self) {
self.has_failed |= other.has_failed;
self.managed_revision_failed |= other.managed_revision_failed;
}
}
struct ReconcileRound {
req: HeartbeatRequest,
machine_id: uuid::Uuid,
user_id: i32,
running_inst_ids: HashSet<String>,
local_configs: Vec<crate::db::entity::user_running_network_configs::Model>,
target_config_revision: Option<String>,
should_apply_runtime_revision: bool,
}
async fn prepare_reconcile_round(
session_data: &std::sync::Weak<RwLock<SessionData>>,
storage: &StorageInner,
rpc_client: &mut SessionRpcClient,
req: HeartbeatRequest,
) -> RoundStatus<ReconcileRound> {
let Some(machine_id) = req.machine_id.map(uuid::Uuid::from) else {
tracing::warn!(?req, "Machine id is not set, ignore");
return RoundStatus::Skip;
};
if !SessionRpcService::runtime_heartbeat_is_current(session_data, &req).await {
tracing::debug!(?machine_id, "skip stale heartbeat reconcile request");
return RoundStatus::Skip;
}
let user_id = match storage
.db
.get_user_id_by_token(req.user_token.clone())
.await
{
Ok(Some(user_id)) => user_id,
Ok(None) => {
tracing::info!("User not found by token: {:?}", req.user_token);
return RoundStatus::Stop;
}
Err(e) => {
tracing::error!("Failed to get user id by token, error: {:?}", e);
return RoundStatus::Stop;
}
};
let applied_config_revision = {
let Some(data) = session_data.upgrade() else {
return RoundStatus::Stop;
};
data.read().await.applied_config_revision.clone()
};
let target_config_revision = match storage
.db
.get_managed_config_revision((user_id, machine_id))
.await
{
Ok(revision) => revision,
Err(e) => {
tracing::error!("Failed to read managed config revision, error: {:?}", e);
return RoundStatus::Stop;
}
};
let should_apply_runtime_revision =
target_config_revision.is_some() && target_config_revision != applied_config_revision;
let running_inst_ids = match running_instance_ids_for_round(
rpc_client,
&req,
user_id,
machine_id,
should_apply_runtime_revision,
)
.await
{
RoundStatus::Ready(ids) => ids,
RoundStatus::Skip => return RoundStatus::Skip,
RoundStatus::Stop => return RoundStatus::Stop,
};
let local_configs = match storage
.db
.list_network_configs((user_id, machine_id), ListNetworkProps::EnabledOnly)
.await
{
Ok(configs) => configs,
Err(e) => {
tracing::error!("Failed to list network configs, error: {:?}", e);
return RoundStatus::Stop;
}
};
RoundStatus::Ready(ReconcileRound {
req,
machine_id,
user_id,
running_inst_ids,
local_configs,
target_config_revision,
should_apply_runtime_revision,
})
}
async fn running_instance_ids_for_round(
rpc_client: &mut SessionRpcClient,
req: &HeartbeatRequest,
user_id: i32,
machine_id: uuid::Uuid,
should_apply_runtime_revision: bool,
) -> RoundStatus<HashSet<String>> {
if !should_apply_runtime_revision {
return RoundStatus::Ready(
req.running_network_instances
.iter()
.map(|x| x.to_string())
.collect(),
);
}
match rpc_client
.list_network_instance(BaseController::default(), ListNetworkInstanceRequest {})
.await
{
Ok(resp) => RoundStatus::Ready(resp.inst_ids.iter().map(|x| x.to_string()).collect()),
Err(error) => {
tracing::warn!(
?user_id,
?machine_id,
?error,
"Failed to refresh running instances for managed config revision"
);
RoundStatus::Skip
}
}
}
async fn sync_running_sources_for_round(
rpc_client: &mut SessionRpcClient,
storage: &StorageInner,
round: &mut ReconcileRound,
) -> RoundStatus<Option<Vec<NetworkMeta>>> {
if !round.req.support_config_source {
return RoundStatus::Ready(None);
}
let ret = if round.running_inst_ids.is_empty() {
Ok(Vec::new())
} else {
rpc_client
.list_network_instance_meta(
BaseController::default(),
ListNetworkInstanceMetaRequest {
inst_ids: managed_config::parse_instance_ids(
round.running_inst_ids.iter().cloned(),
),
},
)
.await
.map(|resp| resp.metas)
};
match ret {
Ok(metas) => {
if let Err(e) = managed_config::sync_running_config_sources(
&storage.db,
round.user_id,
round.machine_id,
&round.local_configs,
&metas,
)
.await
{
tracing::warn!(
user_id = ?round.user_id,
machine_id = ?round.machine_id,
%e,
"Failed to sync running network config sources"
);
} else if !metas.is_empty() {
round.local_configs = match storage
.db
.list_network_configs(
(round.user_id, round.machine_id),
ListNetworkProps::EnabledOnly,
)
.await
{
Ok(configs) => configs,
Err(e) => {
tracing::error!(
"Failed to reload network configs after source sync, error: {:?}",
e
);
return RoundStatus::Stop;
}
};
}
RoundStatus::Ready(Some(metas))
}
Err(e) => {
tracing::warn!(
user_id = ?round.user_id,
%e,
"Failed to list running network instance metadata"
);
RoundStatus::Ready(None)
}
}
}
async fn cleanup_stale_web_source_instances(
session_data: &std::sync::Weak<RwLock<SessionData>>,
storage: &StorageInner,
rpc_client: &mut SessionRpcClient,
round: &ReconcileRound,
running_metas: Option<&[NetworkMeta]>,
desired_web_inst_ids: &HashSet<String>,
cache: &mut ReconcileCache,
) -> RoundStatus<ReconcileOutcome> {
let desired_changed = cache
.last_desired_web_inst_ids
.as_ref()
.is_none_or(|last| last != desired_web_inst_ids);
if cache.cleaned_web_source_instances && !desired_changed {
return RoundStatus::Ready(ReconcileOutcome::default());
}
let db_web_inst_ids = match storage
.db
.list_network_configs((round.user_id, round.machine_id), ListNetworkProps::All)
.await
{
Ok(configs) => managed_config::desired_web_source_instance_ids(&configs),
Err(e) => {
tracing::error!("Failed to list all network configs, error: {:?}", e);
return RoundStatus::Stop;
}
};
let running_web_inst_ids = managed_config::running_web_source_instance_ids(
&round.running_inst_ids,
&db_web_inst_ids,
running_metas,
);
let should_delete_inst_ids = running_web_inst_ids
.difference(desired_web_inst_ids)
.cloned()
.collect::<HashSet<_>>();
let should_delete_ids =
managed_config::parse_instance_ids(should_delete_inst_ids.iter().cloned());
let mut outcome = ReconcileOutcome::default();
if !should_delete_ids.is_empty() {
if !SessionRpcService::runtime_heartbeat_is_current(session_data, &round.req).await {
tracing::debug!(
machine_id = ?round.machine_id,
"skip stale cleanup because webhook session is no longer current"
);
return RoundStatus::Skip;
}
let ret = rpc_client
.delete_network_instance(
BaseController::default(),
DeleteNetworkInstanceRequest {
inst_ids: should_delete_ids,
},
)
.await;
tracing::info!(
user_id = ?round.user_id,
"Clean stale web-source network instances on heartbeat: {:?}, user_token: {:?}",
ret,
round.req.user_token
);
if ret.is_err() {
outcome.record_failure(true);
} else {
cache.runtime_configs.forget_many(&should_delete_inst_ids);
}
}
if !outcome.has_failed {
cache.cleaned_web_source_instances = true;
cache.last_desired_web_inst_ids = Some(desired_web_inst_ids.clone());
}
RoundStatus::Ready(outcome)
}
async fn reconcile_desired_runtime_configs(
session_data: &std::sync::Weak<RwLock<SessionData>>,
rpc_client: &mut SessionRpcClient,
config_client: &mut SessionConfigClient,
round: &ReconcileRound,
cache: &mut ReconcileCache,
) -> ReconcileOutcome {
let mut outcome = ReconcileOutcome::default();
// After stale web-owned instances are removed, start every enabled
// config that the latest heartbeat did not report as running. When
// a managed config revision is pending, also reconcile running
// web-owned configs before reporting that revision as applied.
for config in &round.local_configs {
let source = PersistedConfigSource::from_db(&config.source);
let is_running = round.running_inst_ids.contains(&config.network_instance_id);
let should_reconcile_running_web_config = is_running
&& round.should_apply_runtime_revision
&& source == PersistedConfigSource::Web;
if is_running && !should_reconcile_running_web_config {
continue;
}
let desired_config = match serde_json::from_str::<NetworkConfig>(&config.network_config) {
Ok(cfg) => cfg,
Err(e) => {
tracing::error!(
user_id = ?round.user_id,
machine_id = ?round.machine_id,
instance_id = %config.network_instance_id,
"Failed to deserialize network config, skipping: {:?}",
e
);
if source == PersistedConfigSource::Web {
cache.runtime_configs.forget(&config.network_instance_id);
}
outcome.record_failure(source == PersistedConfigSource::Web);
continue;
}
};
let action_result = if should_reconcile_running_web_config {
reconcile_running_web_config(
session_data,
rpc_client,
config_client,
round,
config,
desired_config,
&mut cache.runtime_configs,
)
.await
} else {
if source == PersistedConfigSource::Web {
cache.runtime_configs.forget(&config.network_instance_id);
}
let action_result = run_missing_network_config(
session_data,
rpc_client,
round,
config,
desired_config.clone(),
)
.await;
if matches!(action_result, ConfigActionResult::Success)
&& source == PersistedConfigSource::Web
{
if let Err(e) = remember_web_runtime_config_after_run(
rpc_client,
&config.network_instance_id,
&desired_config,
&mut cache.runtime_configs,
)
.await
{
tracing::error!(
user_id = ?round.user_id,
machine_id = ?round.machine_id,
instance_id = %config.network_instance_id,
"Failed to cache runtime config after run: {:?}",
e
);
ConfigActionResult::Failed
} else {
action_result
}
} else {
action_result
}
};
match action_result {
ConfigActionResult::Success => {}
ConfigActionResult::Failed => {
if source == PersistedConfigSource::Web {
cache.runtime_configs.forget(&config.network_instance_id);
}
outcome.record_failure(source == PersistedConfigSource::Web)
}
ConfigActionResult::StopRound => {
if source == PersistedConfigSource::Web {
cache.runtime_configs.forget(&config.network_instance_id);
}
outcome.record_failure(source == PersistedConfigSource::Web);
break;
}
}
}
outcome
}
async fn reconcile_running_web_config(
session_data: &std::sync::Weak<RwLock<SessionData>>,
rpc_client: &mut SessionRpcClient,
config_client: &mut SessionConfigClient,
round: &ReconcileRound,
config: &crate::db::entity::user_running_network_configs::Model,
desired_config: NetworkConfig,
runtime_config_cache: &mut SessionRuntimeConfigCache,
) -> ConfigActionResult {
if !SessionRpcService::runtime_heartbeat_is_current(session_data, &round.req).await {
tracing::debug!(
machine_id = ?round.machine_id,
instance_id = %config.network_instance_id,
"skip runtime reconcile because webhook session is no longer current"
);
return ConfigActionResult::StopRound;
}
let ret = async {
let action =
match runtime_config_cache.plan(&config.network_instance_id, desired_config.clone())? {
Some(action) => action,
None => {
runtime_reconcile::prepare_web_source_runtime_reconcile(
&mut *rpc_client,
&config.network_instance_id,
desired_config.clone(),
true,
)
.await?
}
};
if !SessionRpcService::runtime_heartbeat_is_current(session_data, &round.req).await {
anyhow::bail!("webhook session is no longer current before runtime reconcile apply");
}
let observed_config = runtime_reconcile::apply_web_source_runtime_reconcile(
&mut *rpc_client,
&mut *config_client,
&config.network_instance_id,
desired_config.clone(),
action,
)
.await?;
runtime_config_cache.remember(&config.network_instance_id, observed_config);
Ok::<(), anyhow::Error>(())
}
.await;
tracing::info!(
user_id = ?round.user_id,
instance_id = %config.network_instance_id,
"Reconcile running web-source network instance: {:?}, user_token: {:?}",
ret,
round.req.user_token
);
if ret.is_ok() {
ConfigActionResult::Success
} else {
runtime_config_cache.forget(&config.network_instance_id);
ConfigActionResult::Failed
}
}
async fn run_missing_network_config(
session_data: &std::sync::Weak<RwLock<SessionData>>,
rpc_client: &mut SessionRpcClient,
round: &ReconcileRound,
config: &crate::db::entity::user_running_network_configs::Model,
desired_config: NetworkConfig,
) -> ConfigActionResult {
if !SessionRpcService::runtime_heartbeat_is_current(session_data, &round.req).await {
tracing::debug!(
machine_id = ?round.machine_id,
instance_id = %config.network_instance_id,
"skip run network instance because webhook session is no longer current"
);
return ConfigActionResult::StopRound;
}
let ret = rpc_client
.run_network_instance(
BaseController::default(),
RunNetworkInstanceRequest {
inst_id: Some(config.network_instance_id.clone().into()),
config: Some(desired_config),
overwrite: false,
source: PersistedConfigSource::from_db(&config.source).auto_run_rpc_source() as i32,
},
)
.await;
tracing::info!(
user_id = ?round.user_id,
"Run network instance: {:?}, user_token: {:?}",
ret,
round.req.user_token
);
if ret.is_ok() {
ConfigActionResult::Success
} else {
ConfigActionResult::Failed
}
}
async fn remember_web_runtime_config_after_run(
rpc_client: &mut SessionRpcClient,
inst_id: &str,
desired_config: &NetworkConfig,
runtime_config_cache: &mut SessionRuntimeConfigCache,
) -> anyhow::Result<()> {
let observed_config = runtime_reconcile::get_runtime_config(rpc_client, inst_id).await?;
remember_if_runtime_matches_desired(
inst_id,
desired_config,
observed_config,
runtime_config_cache,
)
}
fn remember_if_runtime_matches_desired(
inst_id: &str,
desired_config: &NetworkConfig,
observed_config: NetworkConfig,
runtime_config_cache: &mut SessionRuntimeConfigCache,
) -> anyhow::Result<()> {
let action = runtime_reconcile::prepare_web_source_runtime_reconcile_from_current(
&observed_config,
desired_config.clone(),
)?;
if !matches!(action, runtime_reconcile::RuntimeReconcileAction::None) {
anyhow::bail!("runtime config still differs after managed run");
}
runtime_config_cache.remember(inst_id, observed_config);
Ok(())
}
async fn mark_config_revision_applied_if_current(
session_data: &std::sync::Weak<RwLock<SessionData>>,
storage: &StorageInner,
round: &ReconcileRound,
outcome: &ReconcileOutcome,
) -> RoundStatus<()> {
if outcome.managed_revision_failed || !round.should_apply_runtime_revision {
return RoundStatus::Ready(());
}
let current_target_config_revision = match storage
.db
.get_managed_config_revision((round.user_id, round.machine_id))
.await
{
Ok(revision) => revision,
Err(e) => {
tracing::error!("Failed to verify managed config revision, error: {:?}", e);
return RoundStatus::Stop;
}
};
if current_target_config_revision != round.target_config_revision {
return RoundStatus::Ready(());
}
let Some(data) = session_data.upgrade() else {
return RoundStatus::Stop;
};
let mut data = data.write().await;
if !SessionRpcService::runtime_heartbeat_is_current_locked(&data, &round.req) {
return RoundStatus::Ready(());
}
data.applied_config_revision = round.target_config_revision.clone();
RoundStatus::Ready(())
}
#[cfg(test)]
mod tests {
use easytier::proto::api::manage::{NetworkingMethod, PortForwardConfig};
use super::*;
fn config_with_port_forwards(port_forwards: Vec<PortForwardConfig>) -> NetworkConfig {
NetworkConfig {
instance_id: Some("11111111-1111-1111-1111-111111111111".to_string()),
dhcp: Some(true),
network_name: Some("managed".to_string()),
network_secret: Some("secret".to_string()),
networking_method: Some(NetworkingMethod::Manual as i32),
port_forwards,
..Default::default()
}
}
fn port_forward(bind_port: u32, dst_port: u32) -> PortForwardConfig {
PortForwardConfig {
bind_ip: "127.0.0.1".to_string(),
bind_port,
dst_ip: "10.144.0.1".to_string(),
dst_port,
proto: "tcp".to_string(),
}
}
#[test]
fn session_runtime_config_cache_misses_unknown_instance() {
let cache = SessionRuntimeConfigCache::default();
let action = cache
.plan("missing", config_with_port_forwards(Vec::new()))
.expect("prepare action");
assert!(action.is_none());
}
#[test]
fn session_runtime_config_cache_skips_matching_observed_config() {
let mut cache = SessionRuntimeConfigCache::default();
let config = config_with_port_forwards(vec![port_forward(23000, 5174)]);
cache.remember("managed", config.clone());
let action = cache
.plan("managed", config)
.expect("prepare action")
.expect("cached action");
assert!(matches!(
action,
runtime_reconcile::RuntimeReconcileAction::None
));
}
#[test]
fn session_runtime_config_cache_plans_patch_from_observed_config() {
let mut cache = SessionRuntimeConfigCache::default();
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
cache.remember("managed", current);
let action = cache
.plan("managed", desired)
.expect("prepare action")
.expect("cached action");
let runtime_reconcile::RuntimeReconcileAction::Patch(patch) = action else {
panic!("expected cached runtime config to produce hot patch");
};
assert_eq!(patch.port_forwards.len(), 1);
}
#[test]
fn session_runtime_config_cache_retain_desired_removes_stale_entries() {
let mut cache = SessionRuntimeConfigCache::default();
let config = config_with_port_forwards(Vec::new());
cache.remember("keep", config.clone());
cache.remember("drop", config);
cache.retain_desired(&HashSet::from(["keep".to_string()]));
assert!(cache.entries.contains_key("keep"));
assert!(!cache.entries.contains_key("drop"));
}
#[test]
fn session_runtime_config_cache_forget_removes_observed_config() {
let mut cache = SessionRuntimeConfigCache::default();
let config = config_with_port_forwards(Vec::new());
cache.remember("managed", config.clone());
cache.forget("managed");
let action = cache
.plan("managed", config)
.expect("prepare action after remove");
assert!(action.is_none());
}
#[test]
fn missing_run_remembers_observed_config_when_it_matches_desired() {
let mut cache = SessionRuntimeConfigCache::default();
let config = config_with_port_forwards(vec![port_forward(23000, 5174)]);
remember_if_runtime_matches_desired("managed", &config, config.clone(), &mut cache)
.expect("remember observed config after run");
let action = cache
.plan("managed", config)
.expect("prepare action after run")
.expect("cached action");
assert!(matches!(
action,
runtime_reconcile::RuntimeReconcileAction::None
));
}
#[test]
fn missing_run_does_not_remember_observed_config_that_still_differs() {
let mut cache = SessionRuntimeConfigCache::default();
let current = config_with_port_forwards(vec![port_forward(23000, 5174)]);
let desired =
config_with_port_forwards(vec![port_forward(23000, 5174), port_forward(23007, 3389)]);
let err = remember_if_runtime_matches_desired("managed", &desired, current, &mut cache)
.expect_err("expected stale run result not to be cached");
assert!(
err.to_string()
.contains("runtime config still differs after managed run")
);
let action = cache
.plan("managed", desired)
.expect("prepare action after stale run result");
assert!(action.is_none());
}
}
@@ -1,402 +0,0 @@
use std::{sync::Arc, time::Duration};
use anyhow::Context as _;
use easytier::proto::web::HeartbeatRequest;
use tokio::sync::RwLock;
use super::{
SessionAuthState, SessionData, SessionRpcService, WebhookConnectNotification,
WebhookDisconnectNotification, send_webhook_connection_transition,
};
use crate::{
client_manager::storage::{Storage, StorageToken},
webhook::SharedWebhookConfig,
};
pub(super) const VALIDATION_RETRY_MS: u64 = 60_000;
pub(super) struct WebhookHeartbeatValidation {
pub(super) config_revision: String,
pub(super) binding_version: u64,
}
pub(super) struct WebhookValidationInput {
pub(super) storage: Storage,
pub(super) webhook_config: SharedWebhookConfig,
pub(super) client_url: url::Url,
pub(super) applied_config_revision: Option<String>,
pub(super) req: HeartbeatRequest,
pub(super) machine_id: uuid::Uuid,
}
fn deterministic_machine_delay(machine_id: uuid::Uuid, max_delay_ms: u64) -> Duration {
let delay_ms = (machine_id.as_u128() % u128::from(max_delay_ms + 1)) as u64;
Duration::from_millis(delay_ms)
}
pub(super) fn retry_delay(machine_id: uuid::Uuid) -> Duration {
Duration::from_millis(VALIDATION_RETRY_MS)
+ deterministic_machine_delay(machine_id, VALIDATION_RETRY_MS)
}
async fn request_heartbeat_validation(
webhook_config: &crate::webhook::WebhookConfig,
client_url: &url::Url,
persisted_config_revision: Option<&str>,
applied_config_revision: Option<&str>,
req: &HeartbeatRequest,
machine_id: uuid::Uuid,
) -> anyhow::Result<Option<WebhookHeartbeatValidation>> {
let webhook_req = crate::webhook::ValidateTokenRequest {
token: req.user_token.clone(),
machine_id: machine_id.to_string(),
public_ip: client_url.host_str().map(str::to_string),
hostname: req.hostname.clone(),
version: req.easytier_version.clone(),
os_type: req.device_os.as_ref().map(|info| info.os_type.clone()),
os_version: req.device_os.as_ref().map(|info| info.version.clone()),
os_distribution: req.device_os.as_ref().map(|info| info.distribution.clone()),
web_instance_id: webhook_config.web_instance_id.clone(),
web_instance_api_base_url: webhook_config.web_instance_api_base_url.clone(),
persisted_config_revision: persisted_config_revision.map(str::to_string),
applied_config_revision: applied_config_revision.map(str::to_string),
};
let resp = webhook_config
.validate_token(&webhook_req)
.await
.map_err(|e| anyhow::anyhow!("Webhook token validation failed: {:?}", e))?;
if !resp.valid {
return Ok(None);
}
Ok(Some(WebhookHeartbeatValidation {
config_revision: resp.config_revision,
binding_version: resp.binding_version,
}))
}
async fn resolve_user_id(storage: &Storage, token: &str) -> anyhow::Result<i32> {
let user_id = match storage
.db()
.get_user_id_by_token(token)
.await
.map_err(|e| anyhow::anyhow!("DB error: {:?}", e))?
{
Some(id) => id,
None => storage
.auto_create_user(token)
.await
.with_context(|| format!("Failed to auto-create webhook user: {:?}", token))?,
};
Ok(user_id)
}
async fn persisted_config_revision_for_token(
storage: &Storage,
token: &str,
machine_id: uuid::Uuid,
) -> anyhow::Result<Option<String>> {
let Some(user_id) = storage
.db()
.get_user_id_by_token(token)
.await
.map_err(|e| anyhow::anyhow!("DB error: {:?}", e))?
else {
return Ok(None);
};
storage
.db()
.get_managed_config_revision((user_id, machine_id))
.await
.map_err(|e| anyhow::anyhow!("DB error: {:?}", e))
}
async fn wait_for_input(
session_data: std::sync::Weak<RwLock<SessionData>>,
) -> Option<WebhookValidationInput> {
loop {
let notify = {
let session_data = session_data.upgrade()?;
let mut data = session_data.write().await;
if matches!(data.auth_state, SessionAuthState::Invalid) {
data.webhook_validation_dirty = false;
tracing::info!(
client_url = %data.client_url,
"webhook validation stopped for invalid session; reconnect is required before revalidation"
);
return None;
}
if data.webhook_validation_dirty {
data.webhook_validation_dirty = false;
let req = data.req.clone()?;
let machine_id = req.machine_id.map(Into::into)?;
let storage = Storage::try_from(data.storage.clone()).ok()?;
return Some(WebhookValidationInput {
storage,
webhook_config: data.webhook_config.clone(),
client_url: data.client_url.clone(),
applied_config_revision: data.applied_config_revision.clone(),
req,
machine_id,
});
}
data.webhook_validation_notify.clone()
};
notify.notified().await;
}
}
pub(super) async fn run_worker(session_data: std::sync::Weak<RwLock<SessionData>>) {
while let Some(input) = wait_for_input(session_data.clone()).await {
let machine_id = input.machine_id;
if let Err(error) = run_round(session_data.clone(), input).await {
tracing::warn!(
?machine_id,
%error,
"webhook validation failed, will retry later"
);
tokio::time::sleep(retry_delay(machine_id)).await;
mark_dirty_if_current(&session_data, machine_id).await;
}
}
}
pub(super) async fn run_round(
session_data: std::sync::Weak<RwLock<SessionData>>,
input: WebhookValidationInput,
) -> anyhow::Result<()> {
let persisted_config_revision = persisted_config_revision_for_token(
&input.storage,
&input.req.user_token,
input.machine_id,
)
.await?;
let validation = request_heartbeat_validation(
&input.webhook_config,
&input.client_url,
persisted_config_revision.as_deref(),
input.applied_config_revision.as_deref(),
&input.req,
input.machine_id,
)
.await?;
let Some(validation) = validation else {
apply_rejected(&session_data, &input).await;
return Ok(());
};
let user_id = resolve_user_id(&input.storage, &input.req.user_token).await?;
apply_success(&session_data, input, validation, user_id).await;
Ok(())
}
async fn mark_dirty_if_current(
session_data: &std::sync::Weak<RwLock<SessionData>>,
machine_id: uuid::Uuid,
) {
let Some(session_data) = session_data.upgrade() else {
return;
};
let notify = {
let mut data = session_data.write().await;
let Some(req) = data.req.as_ref() else {
return;
};
if req.machine_id.map(uuid::Uuid::from) != Some(machine_id) {
return;
}
if matches!(data.auth_state, SessionAuthState::Invalid) {
data.webhook_validation_dirty = false;
tracing::debug!(
%machine_id,
"skip webhook validation retry for invalid session"
);
return;
}
SessionRpcService::mark_webhook_validation_dirty_locked(&mut data)
};
notify.notify_one();
}
pub(super) async fn apply_rejected(
session_data: &std::sync::Weak<RwLock<SessionData>>,
input: &WebhookValidationInput,
) {
let Some(session_data) = session_data.upgrade() else {
return;
};
let (storage_token, disconnect_notification) = {
let mut data = session_data.write().await;
if !data.req.as_ref().is_some_and(|req| {
SessionRpcService::heartbeat_matches_identity(
req,
&input.req.user_token,
input.machine_id,
)
}) {
return;
}
tracing::info!(
machine_id = %input.machine_id,
client_url = %data.client_url,
"webhook token rejected; marking session invalid and requiring client reconnect"
);
data.auth_state = SessionAuthState::Invalid;
data.webhook_validation_dirty = false;
data.binding_version = None;
data.applied_config_revision = None;
let storage_token = data.storage_token.clone();
let disconnect_notification = storage_token.as_ref().and_then(|storage_token| {
data.webhook_connected_binding_version
.take()
.map(|binding_version| WebhookDisconnectNotification {
webhook: data.webhook_config.clone(),
storage_token: storage_token.clone(),
binding_version,
})
});
(storage_token, disconnect_notification)
};
if let Some(storage_token) = storage_token {
let report_time = SessionRpcService::heartbeat_report_timestamp(&input.req);
input
.storage
.update_client(storage_token, report_time, false);
}
if disconnect_notification.is_some() {
wait_webhook_connection_transition(
Arc::downgrade(&session_data),
disconnect_notification,
None,
)
.await;
}
}
pub(super) async fn apply_success(
session_data: &std::sync::Weak<RwLock<SessionData>>,
input: WebhookValidationInput,
validation: WebhookHeartbeatValidation,
user_id: i32,
) {
let WebhookHeartbeatValidation {
config_revision: _,
binding_version,
} = validation;
let Some(session_data) = session_data.upgrade() else {
return;
};
let (storage_token, notifier, disconnect_notification, connect_notification, runtime_req) = {
let mut data = session_data.write().await;
let Some(runtime_req) = data.req.clone() else {
return;
};
if !SessionRpcService::heartbeat_matches_identity(
&runtime_req,
&input.req.user_token,
input.machine_id,
) {
return;
}
if matches!(data.auth_state, SessionAuthState::Invalid) {
tracing::info!(
machine_id = %input.machine_id,
client_url = %data.client_url,
"ignore webhook validation success for invalid session; reconnect is required before revalidation"
);
return;
}
let previous_connected_binding_version = data.webhook_connected_binding_version;
let client_url = data.client_url.clone();
let storage_token = data.storage_token.get_or_insert_with(|| StorageToken {
token: runtime_req.user_token.clone(),
client_url,
machine_id: input.machine_id,
user_id,
});
let storage_token = storage_token.clone();
data.auth_state = SessionAuthState::Authorized;
data.binding_version = Some(binding_version);
let should_notify_connected = previous_connected_binding_version != Some(binding_version);
let disconnect_notification = previous_connected_binding_version
.filter(|previous_binding_version| *previous_binding_version != binding_version)
.map(|previous_binding_version| {
data.webhook_connected_binding_version = None;
WebhookDisconnectNotification {
webhook: data.webhook_config.clone(),
storage_token: storage_token.clone(),
binding_version: previous_binding_version,
}
});
let connect_notification = should_notify_connected.then(|| WebhookConnectNotification {
webhook: data.webhook_config.clone(),
storage_token: storage_token.clone(),
binding_version,
req: crate::webhook::NodeConnectedRequest {
machine_id: input.machine_id.to_string(),
token: runtime_req.user_token.clone(),
user_id: Some(user_id),
hostname: runtime_req.hostname.clone(),
version: runtime_req.easytier_version.clone(),
os_type: runtime_req
.device_os
.as_ref()
.map(|info| info.os_type.clone()),
os_version: runtime_req
.device_os
.as_ref()
.map(|info| info.version.clone()),
os_distribution: runtime_req
.device_os
.as_ref()
.map(|info| info.distribution.clone()),
web_instance_id: data.webhook_config.web_instance_id.clone(),
binding_version: Some(binding_version),
},
});
(
storage_token,
data.notifier.clone(),
disconnect_notification,
connect_notification,
runtime_req,
)
};
if disconnect_notification.is_some() || connect_notification.is_some() {
wait_webhook_connection_transition(
Arc::downgrade(&session_data),
disconnect_notification,
connect_notification,
)
.await;
}
let report_time = SessionRpcService::heartbeat_report_timestamp(&runtime_req);
input
.storage
.update_client(storage_token, report_time, true);
let _ = notifier.send(runtime_req);
}
async fn wait_webhook_connection_transition(
session_data: std::sync::Weak<RwLock<SessionData>>,
disconnect: Option<WebhookDisconnectNotification>,
connect: Option<WebhookConnectNotification>,
) {
let transition = tokio::spawn(send_webhook_connection_transition(
session_data,
disconnect,
connect,
));
if let Err(error) = transition.await {
tracing::warn!(%error, "webhook connection transition task failed");
}
}
+9 -123
View File
@@ -17,7 +17,6 @@ pub struct StorageToken {
struct ClientInfo {
storage_token: StorageToken,
report_time: i64,
authorized: bool,
}
#[derive(Debug)]
@@ -56,19 +55,7 @@ impl Storage {
fn update_client_info_map(map: &DashMap<uuid::Uuid, ClientInfo>, client_info: &ClientInfo) {
map.entry(client_info.storage_token.machine_id)
.and_modify(|e| {
let same_client = e.storage_token.client_url
== client_info.storage_token.client_url
&& e.storage_token.user_id == client_info.storage_token.user_id;
let should_replace = if (same_client && e.authorized != client_info.authorized)
|| (!e.authorized && client_info.authorized)
{
true
} else if e.authorized && !client_info.authorized && !same_client {
false
} else {
e.report_time < client_info.report_time
};
if should_replace {
if e.report_time < client_info.report_time {
assert_eq!(
e.storage_token.machine_id,
client_info.storage_token.machine_id
@@ -79,13 +66,12 @@ impl Storage {
.or_insert(client_info.clone());
}
pub fn update_client(&self, stoken: StorageToken, report_time: i64, authorized: bool) {
pub fn update_client(&self, stoken: StorageToken, report_time: i64) {
let inner = self.0.user_clients_map.entry(stoken.user_id).or_default();
let client_info = ClientInfo {
storage_token: stoken.clone(),
report_time,
authorized,
};
Self::update_client_info_map(&inner, &client_info);
}
@@ -107,21 +93,11 @@ impl Storage {
&self,
user_id: UserIdInDb,
machine_id: &uuid::Uuid,
) -> Option<url::Url> {
self.get_client_url_by_machine_id_with_auth(user_id, machine_id, true)
}
pub fn get_client_url_by_machine_id_with_auth(
&self,
user_id: UserIdInDb,
machine_id: &uuid::Uuid,
require_authorized: bool,
) -> Option<url::Url> {
self.0.user_clients_map.get(&user_id).and_then(|info_map| {
info_map.get(machine_id).and_then(|info| {
(!require_authorized || info.authorized)
.then(|| info.storage_token.client_url.clone())
})
info_map
.get(machine_id)
.map(|info| info.storage_token.client_url.clone())
})
}
@@ -132,7 +108,6 @@ impl Storage {
.map(|info_map| {
info_map
.iter()
.filter(|info| info.value().authorized)
.map(|info| info.value().storage_token.client_url.clone())
.collect()
})
@@ -140,14 +115,6 @@ impl Storage {
}
pub fn list_clients(&self) -> Vec<StorageToken> {
self.list_clients_with_auth(true)
}
pub fn list_all_clients(&self) -> Vec<StorageToken> {
self.list_clients_with_auth(false)
}
fn list_clients_with_auth(&self, require_authorized: bool) -> Vec<StorageToken> {
self.0
.user_clients_map
.iter()
@@ -155,7 +122,6 @@ impl Storage {
user_clients
.value()
.iter()
.filter(|info| !require_authorized || info.value().authorized)
.map(|info| info.value().storage_token.clone())
.collect::<Vec<_>>()
})
@@ -198,8 +164,8 @@ mod tests {
let user1_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
let user2_token = make_storage_token(2, machine_id, "tcp://127.0.0.1:1002");
storage.update_client(user1_token.clone(), 10, true);
storage.update_client(user2_token.clone(), 20, true);
storage.update_client(user1_token.clone(), 10);
storage.update_client(user2_token.clone(), 20);
assert_eq!(
storage.get_client_url_by_machine_id(1, &machine_id),
@@ -229,8 +195,8 @@ mod tests {
let user1_token = make_storage_token(1, uuid::Uuid::new_v4(), "tcp://127.0.0.1:1001");
let user2_token = make_storage_token(2, uuid::Uuid::new_v4(), "tcp://127.0.0.1:1002");
storage.update_client(user1_token.clone(), 10, true);
storage.update_client(user2_token.clone(), 20, true);
storage.update_client(user1_token.clone(), 10);
storage.update_client(user2_token.clone(), 20);
let tokens = storage.list_clients();
assert_eq!(tokens.len(), 2);
@@ -243,84 +209,4 @@ mod tests {
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].token, user2_token.token);
}
#[tokio::test]
async fn pending_client_is_listed_but_not_authorized_for_machine_lookup() {
let storage = Storage::new(Db::memory_db().await);
let machine_id = uuid::Uuid::new_v4();
let token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
storage.update_client(token.clone(), 10, false);
assert_eq!(storage.list_clients().len(), 0);
assert_eq!(storage.list_all_clients().len(), 1);
assert_eq!(storage.list_user_clients(1), Vec::<url::Url>::new());
assert_eq!(storage.get_client_url_by_machine_id(1, &machine_id), None);
assert_eq!(
storage.get_client_url_by_machine_id_with_auth(1, &machine_id, false),
Some(token.client_url.clone())
);
storage.update_client(token.clone(), 11, true);
assert_eq!(
storage.get_client_url_by_machine_id(1, &machine_id),
Some(token.client_url.clone())
);
storage.update_client(token.clone(), 11, false);
assert_eq!(storage.get_client_url_by_machine_id(1, &machine_id), None);
assert_eq!(storage.list_clients().len(), 0);
assert_eq!(storage.list_all_clients().len(), 1);
}
#[tokio::test]
async fn stale_client_authorization_update_does_not_replace_newer_client() {
let storage = Storage::new(Db::memory_db().await);
let machine_id = uuid::Uuid::new_v4();
let old_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
let new_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1002");
storage.update_client(old_token.clone(), 10, true);
storage.update_client(new_token.clone(), 20, true);
storage.update_client(old_token, 10, false);
assert_eq!(
storage.get_client_url_by_machine_id(1, &machine_id),
Some(new_token.client_url)
);
}
#[tokio::test]
async fn pending_client_does_not_replace_authorized_route() {
let storage = Storage::new(Db::memory_db().await);
let machine_id = uuid::Uuid::new_v4();
let authorized_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
let pending_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1002");
storage.update_client(authorized_token.clone(), 10, true);
storage.update_client(pending_token, i64::MAX, false);
assert_eq!(
storage.get_client_url_by_machine_id(1, &machine_id),
Some(authorized_token.client_url)
);
}
#[tokio::test]
async fn authorized_client_replaces_pending_route_regardless_of_report_time() {
let storage = Storage::new(Db::memory_db().await);
let machine_id = uuid::Uuid::new_v4();
let pending_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
let authorized_token = make_storage_token(1, machine_id, "tcp://127.0.0.1:1002");
storage.update_client(pending_token, i64::MAX, false);
storage.update_client(authorized_token.clone(), 10, true);
assert_eq!(
storage.get_client_url_by_machine_id(1, &machine_id),
Some(authorized_token.client_url)
);
}
}
@@ -1,38 +0,0 @@
//! `SeaORM` Entity, hand-written to match the generated entity style.
use sea_orm::entity::prelude::*;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Eq, Serialize, Deserialize)]
#[sea_orm(table_name = "managed_config_revisions")]
pub struct Model {
#[sea_orm(primary_key)]
pub id: i32,
pub user_id: i32,
#[sea_orm(column_type = "Text")]
pub device_id: String,
#[sea_orm(column_type = "Text")]
pub config_revision: String,
pub create_time: DateTimeWithTimeZone,
pub update_time: DateTimeWithTimeZone,
}
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
pub enum Relation {
#[sea_orm(
belongs_to = "super::users::Entity",
from = "Column::UserId",
to = "super::users::Column::Id",
on_update = "Cascade",
on_delete = "Cascade"
)]
Users,
}
impl Related<super::users::Entity> for Entity {
fn to() -> RelationDef {
Relation::Users.def()
}
}
impl ActiveModelBehavior for ActiveModel {}
-1
View File
@@ -4,7 +4,6 @@ pub mod prelude;
pub mod groups;
pub mod groups_permissions;
pub mod managed_config_revisions;
pub mod permissions;
pub mod tower_sessions;
pub mod user_running_network_configs;
-1
View File
@@ -2,7 +2,6 @@
pub use super::groups::Entity as Groups;
pub use super::groups_permissions::Entity as GroupsPermissions;
pub use super::managed_config_revisions::Entity as ManagedConfigRevisions;
pub use super::permissions::Entity as Permissions;
pub use super::tower_sessions::Entity as TowerSessions;
pub use super::user_running_network_configs::Entity as UserRunningNetworkConfigs;
-146
View File
@@ -141,110 +141,6 @@ impl Db {
) -> Result<Option<UserIdInDb>, DbErr> {
self.get_user_id(token).await
}
pub async fn get_managed_config_revision(
&self,
(user_id, device_id): (UserIdInDb, Uuid),
) -> Result<Option<String>, DbErr> {
use entity::managed_config_revisions as mcr;
let revision = mcr::Entity::find()
.filter(mcr::Column::UserId.eq(user_id))
.filter(mcr::Column::DeviceId.eq(device_id.to_string()))
.one(self.orm_db())
.await?;
Ok(revision.map(|row| row.config_revision))
}
pub async fn set_managed_config_revision(
&self,
(user_id, device_id): (UserIdInDb, Uuid),
config_revision: &str,
) -> Result<(), DbErr> {
use entity::managed_config_revisions as mcr;
let now = chrono::Local::now().fixed_offset();
let on_conflict = OnConflict::columns([mcr::Column::UserId, mcr::Column::DeviceId])
.update_columns([mcr::Column::ConfigRevision, mcr::Column::UpdateTime])
.to_owned();
let insert_m = mcr::ActiveModel {
user_id: Set(user_id),
device_id: Set(device_id.to_string()),
config_revision: Set(config_revision.to_string()),
create_time: Set(now),
update_time: Set(now),
..Default::default()
};
mcr::Entity::insert(insert_m)
.on_conflict(on_conflict)
.do_nothing()
.exec(self.orm_db())
.await?;
Ok(())
}
pub async fn insert_or_update_web_network_config(
&self,
(user_id, device_id): (UserIdInDb, Uuid),
network_inst_id: Uuid,
network_config: NetworkConfig,
) -> Result<bool, DbErr> {
let now = chrono::Local::now().fixed_offset();
let network_config =
serde_json::to_string(&network_config).map_err(|e| DbErr::Json(e.to_string()))?;
let source = ConfigSource::Web.as_str();
let result = sqlx::query(
r#"
INSERT INTO user_running_network_configs (
user_id, device_id, network_instance_id, network_config,
source, disabled, create_time, update_time
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(user_id, device_id, network_instance_id) DO UPDATE SET
network_config = excluded.network_config,
source = excluded.source,
disabled = excluded.disabled,
update_time = excluded.update_time
WHERE user_running_network_configs.source = ?
"#,
)
.bind(user_id)
.bind(device_id.to_string())
.bind(network_inst_id.to_string())
.bind(network_config)
.bind(source)
.bind(false)
.bind(now)
.bind(now)
.bind(source)
.execute(&self.db)
.await
.map_err(|e| DbErr::Custom(e.to_string()))?;
Ok(result.rows_affected() > 0)
}
pub async fn delete_web_network_configs(
&self,
(user_id, device_id): (UserIdInDb, Uuid),
network_inst_ids: &[Uuid],
) -> Result<(), DbErr> {
use entity::user_running_network_configs as urnc;
urnc::Entity::delete_many()
.filter(urnc::Column::UserId.eq(user_id))
.filter(urnc::Column::DeviceId.eq(device_id.to_string()))
.filter(urnc::Column::Source.eq(ConfigSource::Web.as_str()))
.filter(
urnc::Column::NetworkInstanceId
.is_in(network_inst_ids.iter().map(|id| id.to_string())),
)
.exec(self.orm_db())
.await?;
Ok(())
}
}
#[async_trait]
@@ -572,46 +468,4 @@ mod tests {
assert_eq!(device1_configs.len(), 1);
assert_eq!(device2_configs.len(), 1);
}
#[tokio::test]
async fn test_web_network_config_does_not_replace_user_owned_config() {
let db = Db::memory_db().await;
let user_id = db.auto_create_user("user-web-race").await.unwrap().id;
let device_id = uuid::Uuid::new_v4();
let inst_id = uuid::Uuid::new_v4();
db.insert_or_update_user_network_config(
(user_id, device_id),
inst_id,
NetworkConfig {
network_name: Some("user-owned".to_string()),
..Default::default()
},
ConfigSource::User,
)
.await
.unwrap();
let updated = db
.insert_or_update_web_network_config(
(user_id, device_id),
inst_id,
NetworkConfig {
network_name: Some("web-owned".to_string()),
..Default::default()
},
)
.await
.unwrap();
assert!(!updated);
let saved = db
.get_network_config((user_id, device_id), &inst_id.to_string())
.await
.unwrap()
.unwrap();
assert_eq!(saved.get_network_config_source(), ConfigSource::User);
let saved_config = saved.get_network_config().unwrap();
assert_eq!(saved_config.network_name.as_deref(), Some("user-owned"));
}
}
+1 -10
View File
@@ -3,8 +3,8 @@
#[macro_use]
extern crate rust_i18n;
use std::net::IpAddr;
use std::sync::Arc;
use std::{net::IpAddr, time::Duration};
use clap::Parser;
use easytier::tunnel::websocket::WsTunnelListener;
@@ -113,14 +113,6 @@ struct Cli {
)]
geoip_db: Option<String>,
#[arg(
long,
env = "ET_HEARTBEAT_MIN_RESPONSE_MS",
default_value = "0",
help = t!("cli.heartbeat_min_response_ms").to_string(),
)]
heartbeat_min_response_ms: u64,
#[cfg(feature = "embed")]
#[arg(
long,
@@ -320,7 +312,6 @@ async fn main() {
let mut mgr = client_manager::ClientManager::new(
db.clone(),
cli.geoip_db,
Duration::from_millis(cli.heartbeat_min_response_ms),
feature_flags.clone(),
webhook_config.clone(),
);
@@ -1,46 +0,0 @@
use sea_orm_migration::prelude::*;
pub struct Migration;
impl MigrationName for Migration {
fn name(&self) -> &str {
"m20260619_000005_managed_config_revisions"
}
}
#[async_trait::async_trait]
impl MigrationTrait for Migration {
async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.get_connection()
.execute_unprepared(
r#"
CREATE TABLE managed_config_revisions (
id INTEGER PRIMARY KEY AUTOINCREMENT NOT NULL,
user_id INTEGER NOT NULL,
device_id TEXT NOT NULL,
config_revision TEXT NOT NULL,
create_time TEXT NOT NULL,
update_time TEXT NOT NULL,
CONSTRAINT fk_managed_config_revisions_user_id_to_users_id
FOREIGN KEY (user_id) REFERENCES users(id)
ON DELETE CASCADE
ON UPDATE CASCADE
);
CREATE UNIQUE INDEX idx_managed_config_revisions_scope
ON managed_config_revisions(user_id, device_id);
"#,
)
.await?;
Ok(())
}
async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> {
manager
.get_connection()
.execute_unprepared("DROP TABLE managed_config_revisions;")
.await?;
Ok(())
}
}
-2
View File
@@ -4,7 +4,6 @@ mod m20241029_000001_init;
mod m20260403_000002_scope_network_config_unique;
mod m20260421_000003_add_network_config_source;
mod m20260514_000004_rename_web_config_source;
mod m20260619_000005_managed_config_revisions;
pub struct Migrator;
@@ -16,7 +15,6 @@ impl MigratorTrait for Migrator {
Box::new(m20260403_000002_scope_network_config_unique::Migration),
Box::new(m20260421_000003_add_network_config_source::Migration),
Box::new(m20260514_000004_rename_web_config_source::Migration),
Box::new(m20260619_000005_managed_config_revisions::Migration),
]
}
}
+1 -1
View File
@@ -307,7 +307,7 @@ impl RestfulServer {
async fn handle_list_all_sessions_internal(
State(client_mgr): AppState,
) -> Result<Json<ListSessionJsonResp>, HttpHandleError> {
let ret = client_mgr.list_all_sessions().await;
let ret = client_mgr.list_sessions().await;
Ok(ListSessionJsonResp(ret).into())
}
+5 -15
View File
@@ -93,8 +93,6 @@ struct ManagedNetworkConfigJson {
#[derive(Debug, serde::Deserialize, serde::Serialize)]
struct ReconcileManagedNetworkConfigsJsonReq {
managed_network_configs: Vec<ManagedNetworkConfigJson>,
config_revision: Option<String>,
expected_config_revision: Option<String>,
}
#[derive(Debug, serde::Deserialize, serde::Serialize)]
@@ -359,21 +357,13 @@ impl NetworkApi {
})
.collect();
client_mgr
.reconcile_managed_network_configs(
user_id,
machine_id,
desired,
payload.config_revision,
payload.expected_config_revision,
)
.reconcile_managed_network_configs(user_id, machine_id, desired)
.await
.map_err(|err| {
let status = if crate::client_manager::is_managed_config_revision_conflict(&err) {
StatusCode::CONFLICT
} else {
StatusCode::INTERNAL_SERVER_ERROR
};
(status, other_error(err.to_string()).into())
(
StatusCode::INTERNAL_SERVER_ERROR,
other_error(err.to_string()).into(),
)
})?;
Ok(Void::default().into())
}
+16 -589
View File
@@ -1,248 +1,6 @@
use std::{
cmp::Ordering,
collections::VecDeque,
fmt,
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use std::sync::Arc;
use serde::{Deserialize, Serialize};
use tokio::sync::oneshot;
const VALIDATE_TOKEN_INITIAL_CONCURRENCY: usize = 8;
const VALIDATE_TOKEN_MIN_CONCURRENCY: usize = 2;
const VALIDATE_TOKEN_MAX_CONCURRENCY: usize = 64;
const VALIDATE_TOKEN_ADJUST_WINDOW: Duration = Duration::from_secs(1);
const VALIDATE_TOKEN_SLOW_THRESHOLD: Duration = Duration::from_secs(2);
const WEBHOOK_HTTP_TIMEOUT: Duration = Duration::from_secs(10);
struct AdaptiveValidateLimiter {
state: Mutex<AdaptiveValidateLimiterState>,
}
impl fmt::Debug for AdaptiveValidateLimiter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AdaptiveValidateLimiter")
.field("state", &self.lock_state())
.finish_non_exhaustive()
}
}
struct AdaptiveValidateLimiterState {
limit: usize,
in_flight: usize,
waiters: VecDeque<oneshot::Sender<AdaptiveValidateGrant>>,
window_started_at: Instant,
samples: usize,
slow_samples: usize,
failures: usize,
had_queue: bool,
}
impl fmt::Debug for AdaptiveValidateLimiterState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AdaptiveValidateLimiterState")
.field("limit", &self.limit)
.field("in_flight", &self.in_flight)
.field("waiters", &self.waiters.len())
.field("window_started_at", &self.window_started_at)
.field("samples", &self.samples)
.field("slow_samples", &self.slow_samples)
.field("failures", &self.failures)
.field("had_queue", &self.had_queue)
.finish()
}
}
struct AdaptiveValidatePermit {
limiter: Arc<AdaptiveValidateLimiter>,
started_at: Instant,
completed: bool,
}
struct AdaptiveValidateGrant {
limiter: Arc<AdaptiveValidateLimiter>,
active: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum LimitAdjustment {
Unchanged,
Increased,
Decreased,
}
impl AdaptiveValidateLimiter {
fn new() -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(AdaptiveValidateLimiterState::new(Instant::now())),
})
}
async fn acquire(self: &Arc<Self>) -> AdaptiveValidatePermit {
loop {
let receiver = {
let mut state = self.lock_state();
state.complete_window_if_due(Instant::now());
if state.waiters.is_empty() && state.in_flight < state.limit {
state.in_flight += 1;
return AdaptiveValidatePermit::new(self.clone());
}
let (sender, receiver) = oneshot::channel();
state.had_queue = true;
state.waiters.push_back(sender);
self.grant_waiters(&mut state);
receiver
};
if let Ok(grant) = receiver.await {
return grant.into_permit();
}
}
}
fn grant_waiters(self: &Arc<Self>, state: &mut AdaptiveValidateLimiterState) {
while state.in_flight < state.limit {
let Some(waiter) = state.waiters.pop_front() else {
break;
};
state.in_flight += 1;
if let Err(mut grant) = waiter.send(AdaptiveValidateGrant::new(self.clone())) {
grant.disarm();
state.in_flight -= 1;
}
}
}
fn record_sample(self: &Arc<Self>, elapsed: Duration, success: bool) {
let mut state = self.lock_state();
let adjustment = state.record_sample(Instant::now(), elapsed, success);
if adjustment == LimitAdjustment::Increased {
self.grant_waiters(&mut state);
}
}
fn release_slot(self: &Arc<Self>) {
let mut state = self.lock_state();
state.in_flight = state.in_flight.saturating_sub(1);
self.grant_waiters(&mut state);
}
fn lock_state(&self) -> std::sync::MutexGuard<'_, AdaptiveValidateLimiterState> {
self.state
.lock()
.expect("adaptive validate limiter state should not be poisoned")
}
}
impl AdaptiveValidateLimiterState {
fn new(now: Instant) -> Self {
Self {
limit: VALIDATE_TOKEN_INITIAL_CONCURRENCY,
in_flight: 0,
waiters: VecDeque::new(),
window_started_at: now,
samples: 0,
slow_samples: 0,
failures: 0,
had_queue: false,
}
}
fn record_sample(&mut self, now: Instant, elapsed: Duration, success: bool) -> LimitAdjustment {
self.samples += 1;
if elapsed > VALIDATE_TOKEN_SLOW_THRESHOLD {
self.slow_samples += 1;
}
if !success {
self.failures += 1;
}
self.complete_window_if_due(now)
}
fn complete_window_if_due(&mut self, now: Instant) -> LimitAdjustment {
if now.duration_since(self.window_started_at) < VALIDATE_TOKEN_ADJUST_WINDOW {
return LimitAdjustment::Unchanged;
}
let old_limit = self.limit;
if self.samples > 0 {
if self.failures > 0 || self.is_p95_slow() {
self.limit = (self.limit / 2).max(VALIDATE_TOKEN_MIN_CONCURRENCY);
} else if self.had_queue {
self.limit = (self.limit + 1).min(VALIDATE_TOKEN_MAX_CONCURRENCY);
}
}
self.window_started_at = now;
self.samples = 0;
self.slow_samples = 0;
self.failures = 0;
self.had_queue = false;
match self.limit.cmp(&old_limit) {
Ordering::Greater => LimitAdjustment::Increased,
Ordering::Less => LimitAdjustment::Decreased,
Ordering::Equal => LimitAdjustment::Unchanged,
}
}
fn is_p95_slow(&self) -> bool {
self.slow_samples > 0 && self.slow_samples * 20 >= self.samples
}
}
impl AdaptiveValidatePermit {
fn new(limiter: Arc<AdaptiveValidateLimiter>) -> Self {
Self {
limiter,
started_at: Instant::now(),
completed: false,
}
}
fn complete(mut self, success: bool) {
self.limiter
.record_sample(self.started_at.elapsed(), success);
self.completed = true;
}
}
impl AdaptiveValidateGrant {
fn new(limiter: Arc<AdaptiveValidateLimiter>) -> Self {
Self {
limiter,
active: true,
}
}
fn into_permit(mut self) -> AdaptiveValidatePermit {
self.active = false;
AdaptiveValidatePermit::new(self.limiter.clone())
}
fn disarm(&mut self) {
self.active = false;
}
}
impl Drop for AdaptiveValidateGrant {
fn drop(&mut self) {
if self.active {
self.limiter.release_slot();
}
}
}
impl Drop for AdaptiveValidatePermit {
fn drop(&mut self) {
if !self.completed {
self.limiter.record_sample(self.started_at.elapsed(), false);
}
self.limiter.release_slot();
}
}
/// Webhook configuration for external integrations.
#[derive(Debug, Clone)]
@@ -253,7 +11,6 @@ pub struct WebhookConfig {
pub web_instance_id: Option<String>,
pub web_instance_api_base_url: Option<String>,
validate_limiter: Arc<AdaptiveValidateLimiter>,
client: reqwest::Client,
}
@@ -271,11 +28,7 @@ impl WebhookConfig {
internal_auth_token,
web_instance_id,
web_instance_api_base_url,
validate_limiter: AdaptiveValidateLimiter::new(),
client: reqwest::Client::builder()
.timeout(WEBHOOK_HTTP_TIMEOUT)
.build()
.expect("webhook HTTP client should be valid"),
client: reqwest::Client::new(),
}
}
@@ -305,8 +58,6 @@ pub struct ValidateTokenRequest {
pub web_instance_id: Option<String>,
pub web_instance_api_base_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub persisted_config_revision: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub applied_config_revision: Option<String>,
}
@@ -318,6 +69,7 @@ pub struct ValidateTokenResponse {
#[serde(default)]
pub binding_version: u64,
#[serde(default)]
pub managed_network_configs: Option<Vec<ManagedNetworkConfig>>,
pub config_revision: String,
}
@@ -373,40 +125,21 @@ impl WebhookConfig {
pub async fn validate_token(
&self,
req: &ValidateTokenRequest,
) -> anyhow::Result<ValidateTokenResponse> {
self.validate_token_with_http_timeout(req, WEBHOOK_HTTP_TIMEOUT)
.await
}
async fn validate_token_with_http_timeout(
&self,
req: &ValidateTokenRequest,
http_timeout: Duration,
) -> anyhow::Result<ValidateTokenResponse> {
let url = self.webhook_endpoint("validate-token")?;
let permit = self.validate_limiter.acquire().await;
let ret = match tokio::time::timeout(http_timeout, async {
let resp = self
.client
.post(&url)
.header("X-Internal-Auth", self.webhook_auth_secret())
.json(req)
.send()
.await?;
let resp = self
.client
.post(&url)
.header("X-Internal-Auth", self.webhook_auth_secret())
.json(req)
.send()
.await?;
if !resp.status().is_success() {
anyhow::bail!("webhook validate-token returned status {}", resp.status());
}
if !resp.status().is_success() {
anyhow::bail!("webhook validate-token returned status {}", resp.status());
}
Ok(resp.json().await?)
})
.await
{
Ok(ret) => ret,
Err(_) => Err(anyhow::anyhow!("webhook validate-token timed out")),
};
permit.complete(ret.is_ok());
ret
Ok(resp.json().await?)
}
/// Notify the webhook receiver that a node has connected.
@@ -458,319 +191,13 @@ pub type SharedWebhookConfig = Arc<WebhookConfig>;
#[cfg(test)]
mod tests {
use super::*;
use axum::{Json, Router, routing::post};
use serde_json::json;
#[test]
fn adaptive_validate_limiter_increases_under_queue_pressure() {
let now = Instant::now();
let mut state = AdaptiveValidateLimiterState::new(now);
state.had_queue = true;
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
state.record_sample(now, Duration::from_millis(50), true);
}
assert_eq!(
state.complete_window_if_due(now + VALIDATE_TOKEN_ADJUST_WINDOW),
LimitAdjustment::Increased
);
assert_eq!(state.limit, VALIDATE_TOKEN_INITIAL_CONCURRENCY + 1);
}
#[test]
fn adaptive_validate_limiter_does_not_increase_without_queue_pressure() {
let now = Instant::now();
let mut state = AdaptiveValidateLimiterState::new(now);
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
state.record_sample(now, Duration::from_millis(50), true);
}
assert_eq!(
state.complete_window_if_due(now + VALIDATE_TOKEN_ADJUST_WINDOW),
LimitAdjustment::Unchanged
);
assert_eq!(state.limit, VALIDATE_TOKEN_INITIAL_CONCURRENCY);
}
#[test]
fn adaptive_validate_limiter_reduces_on_failure() {
let now = Instant::now();
let mut state = AdaptiveValidateLimiterState::new(now);
state.record_sample(now, Duration::from_millis(50), false);
assert_eq!(
state.complete_window_if_due(now + VALIDATE_TOKEN_ADJUST_WINDOW),
LimitAdjustment::Decreased
);
assert_eq!(
state.limit,
(VALIDATE_TOKEN_INITIAL_CONCURRENCY / 2).max(VALIDATE_TOKEN_MIN_CONCURRENCY)
);
}
#[test]
fn adaptive_validate_limiter_reduces_on_slow_latency() {
let now = Instant::now();
let mut state = AdaptiveValidateLimiterState::new(now);
state.record_sample(
now,
VALIDATE_TOKEN_SLOW_THRESHOLD + Duration::from_millis(1),
true,
);
assert_eq!(
state.complete_window_if_due(now + VALIDATE_TOKEN_ADJUST_WINDOW),
LimitAdjustment::Decreased
);
assert_eq!(
state.limit,
(VALIDATE_TOKEN_INITIAL_CONCURRENCY / 2).max(VALIDATE_TOKEN_MIN_CONCURRENCY)
);
}
#[tokio::test]
async fn adaptive_validate_limiter_waiter_acquires_after_release() {
let limiter = AdaptiveValidateLimiter::new();
let mut permits = Vec::new();
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
permits.push(limiter.acquire().await);
}
let waiter_limiter = limiter.clone();
let waiter = tokio::spawn(async move {
let permit = waiter_limiter.acquire().await;
permit.complete(true);
});
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!waiter.is_finished());
permits.pop().unwrap().complete(true);
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.unwrap()
.unwrap();
for permit in permits {
permit.complete(true);
}
}
#[tokio::test]
async fn adaptive_validate_limiter_releases_when_permit_is_dropped() {
let limiter = AdaptiveValidateLimiter::new();
let mut permits = Vec::new();
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
permits.push(limiter.acquire().await);
}
let waiter_limiter = limiter.clone();
let waiter = tokio::spawn(async move {
let permit = waiter_limiter.acquire().await;
permit.complete(true);
});
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!waiter.is_finished());
drop(permits.pop().unwrap());
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.unwrap()
.unwrap();
for permit in permits {
permit.complete(true);
}
}
#[tokio::test]
async fn adaptive_validate_limiter_skips_canceled_waiters() {
let limiter = AdaptiveValidateLimiter::new();
let mut permits = Vec::new();
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
permits.push(limiter.acquire().await);
}
let waiter_limiter = limiter.clone();
let waiter = tokio::spawn(async move {
let permit = waiter_limiter.acquire().await;
permit.complete(true);
});
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!waiter.is_finished());
waiter.abort();
assert!(waiter.await.unwrap_err().is_cancelled());
permits.pop().unwrap().complete(true);
tokio::time::sleep(Duration::from_millis(10)).await;
let state = limiter.lock_state();
assert_eq!(state.samples, 1);
assert_eq!(state.failures, 0);
drop(state);
for permit in permits {
permit.complete(true);
}
}
#[test]
fn adaptive_validate_limiter_releases_dropped_grant_without_failure_sample() {
let limiter = AdaptiveValidateLimiter::new();
{
let mut state = limiter.lock_state();
state.in_flight = 1;
}
drop(AdaptiveValidateGrant::new(limiter.clone()));
let state = limiter.lock_state();
assert_eq!(state.in_flight, 0);
assert_eq!(state.samples, 0);
assert_eq!(state.failures, 0);
}
#[tokio::test]
async fn adaptive_validate_limiter_wakes_multiple_waiters_in_order() {
let limiter = AdaptiveValidateLimiter::new();
let mut permits = Vec::new();
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
permits.push(limiter.acquire().await);
}
let (first_acquired_tx, first_acquired_rx) = oneshot::channel();
let (first_release_tx, first_release_rx) = oneshot::channel();
let first = {
let limiter = limiter.clone();
tokio::spawn(async move {
let permit = limiter.acquire().await;
first_acquired_tx.send(()).unwrap();
first_release_rx.await.unwrap();
permit.complete(true);
})
};
let (second_acquired_tx, mut second_acquired_rx) = oneshot::channel();
let second = {
let limiter = limiter.clone();
tokio::spawn(async move {
let permit = limiter.acquire().await;
second_acquired_tx.send(()).unwrap();
permit.complete(true);
})
};
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!first.is_finished());
assert!(!second.is_finished());
permits.pop().unwrap().complete(true);
tokio::time::timeout(Duration::from_secs(1), first_acquired_rx)
.await
.unwrap()
.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(50), &mut second_acquired_rx)
.await
.is_err()
);
first_release_tx.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(1), first)
.await
.unwrap()
.unwrap();
tokio::time::timeout(Duration::from_secs(1), &mut second_acquired_rx)
.await
.unwrap()
.unwrap();
tokio::time::timeout(Duration::from_secs(1), second)
.await
.unwrap()
.unwrap();
for permit in permits {
permit.complete(true);
}
}
#[tokio::test]
async fn validate_token_http_timeout_starts_after_limiter_permit() {
let app = Router::new().route(
"/validate-token",
post(|| async {
Json(json!({
"valid": true,
"config_revision": "rev-1"
}))
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let webhook = WebhookConfig::new(Some(format!("http://{addr}")), None, None, None, None);
let mut permits = Vec::new();
for _ in 0..VALIDATE_TOKEN_INITIAL_CONCURRENCY {
permits.push(webhook.validate_limiter.acquire().await);
}
let validate_webhook = webhook.clone();
let validate = tokio::spawn(async move {
let req = ValidateTokenRequest {
token: "token".to_string(),
machine_id: uuid::Uuid::new_v4().to_string(),
public_ip: None,
hostname: String::new(),
version: String::new(),
os_type: None,
os_version: None,
os_distribution: None,
web_instance_id: None,
web_instance_api_base_url: None,
persisted_config_revision: None,
applied_config_revision: None,
};
validate_webhook
.validate_token_with_http_timeout(&req, Duration::from_millis(20))
.await
});
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!validate.is_finished());
permits.pop().unwrap().complete(true);
let resp = tokio::time::timeout(Duration::from_secs(1), validate)
.await
.unwrap()
.unwrap()
.unwrap();
assert!(resp.valid);
for permit in permits {
permit.complete(true);
}
server.abort();
}
#[test]
fn validate_token_response_deserializes_config_revision() {
fn validate_token_response_allows_missing_managed_configs() {
let resp: ValidateTokenResponse =
serde_json::from_str(r#"{"valid":true,"config_revision":"rev-1"}"#).unwrap();
assert!(resp.valid);
assert_eq!(resp.config_revision, "rev-1");
}
#[test]
fn validate_token_response_allows_missing_config_revision() {
let resp: ValidateTokenResponse = serde_json::from_str(r#"{"valid":true}"#).unwrap();
assert!(resp.valid);
assert!(resp.config_revision.is_empty());
assert!(resp.managed_network_configs.is_none());
}
}
+5 -23
View File
@@ -28,14 +28,6 @@ path = "src/easytier-cli.rs"
name = "easytier"
path = "src/lib.rs"
[[bench]]
name = "tx_throughput"
harness = false
[[bench]]
name = "packet_bytes_extraction"
harness = false
[dependencies]
git-version = "0.3.9"
@@ -59,8 +51,7 @@ time = "0.3"
toml = "0.8.12"
chrono = { version = "0.4.37", features = ["serde"] }
guarden = "0.2"
quanta = "0.12"
guarden = "0.1"
delegate = "0.13.5"
@@ -91,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"
@@ -100,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",
@@ -138,12 +128,11 @@ once_cell = "1.18.0"
# for rpc
prost = "0.14.3"
prost-reflect = { version = "0.16.4", default-features = false, features = ["derive", "serde"] }
prost-reflect = { version = "0.16.4", default-features = false, features = ["derive"] }
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"
@@ -344,7 +333,6 @@ zip = "4.0.0"
[dev-dependencies]
criterion = "0.5.1"
serial_test = "3.0.0"
rstest = "0.25.0"
futures-util = "0.3.31"
@@ -385,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"]
@@ -414,11 +402,5 @@ tracing = ["tokio/tracing", "dep:console-subscriber"]
magic-dns = ["dep:hickory-client", "dep:hickory-server"]
faketcp = ["dep:flume"]
zstd = ["dep:zstd"]
# Deprecated: hotpath profiling has been removed. These feature aliases are
# retained as no-ops so existing build scripts using `--features hotpath*`
# continue to work without pulling in any dependencies.
hotpath = []
hotpath-cpu = ["hotpath"]
hotpath-alloc = ["hotpath"]
# For Network Extension on macOS
macos-ne = []
-162
View File
@@ -1,162 +0,0 @@
# Benchmarks
Criterion benchmarks for EasyTier hot paths.
| Bench | What it measures |
| --------------------------- | -------------------------------------------------------------------------------- |
| `tx_throughput` | End-to-end TX injection path through `peer_manager::send_msg_by_ip` |
| `packet_bytes_extraction` | `ZCPacket::payload_bytes` / `tunnel_payload_bytes` extraction (advance hot path) |
## Packet Bytes Extraction
Criterion benchmark for `ZCPacket` bytes extraction — the methods touched by the
`advance`-based slicing refactor. Measures `payload_bytes` and
`tunnel_payload_bytes` at two payload sizes (1280, 4096). Setup
(`ZCPacket::new_with_payload`) runs in the benchmark harness's preparation
phase and is excluded from the timed region, so the numbers reflect only the
extraction call.
### Quick start
```bash
cargo bench --bench packet_bytes_extraction
```
Smoke run:
```bash
PACKET_BYTES_MEASUREMENT_SECS=2 \
PACKET_BYTES_WARMUP_SECS=1 \
PACKET_BYTES_SAMPLE_SIZE=10 \
cargo bench --bench packet_bytes_extraction -- --quiet
```
### Environment variables
| Variable | Default | Notes |
| ------------------------------- | ------- | ---------------------------- |
| `PACKET_BYTES_MEASUREMENT_SECS` | `10` | Criterion `measurement_time` |
| `PACKET_BYTES_WARMUP_SECS` | `3` | Criterion `warm_up_time` |
| `PACKET_BYTES_SAMPLE_SIZE` | `10` | Criterion `sample_size` (min 10) |
---
## TX Throughput Benchmark
Criterion benchmark for EasyTier's TX injection path (`peer_manager::send_msg_by_ip`).
## What it measures
The benchmark sets up two EasyTier instances (`hot-a` / `hot-b`) and drives
packets from `hot-a` to `hot-b` via `peer_manager.send_msg_by_ip`. This is the
same entry point `easytier-core` uses for daily forwarded traffic, so the
numbers reflect the real TX hot path: NIC pipeline → route lookup →
compress/encrypt → peer connection → tunnel send.
Two variants are reported per tunnel kind:
| Bench | What it measures |
| --------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
| `tx_throughput/<tunnel>` | Serial baseline. One send in flight at a time. Reports per-packet CPU cost (TX injection latency). |
| `tx_throughput/<tunnel>-saturate` | Spawns `TX_THROUGHPUT_INFLIGHT` tokio tasks that independently pump `send_msg_by_ip`. Reports the aggregate throughput ceiling the peer manager + tunnel can sustain across worker threads. |
> **Out of scope (by design):** TUN read/write (`no_tun = true`), compression
> (default `None`), reverse/RX-side measurement, multi-peer fanout. Add
> separate benchmarks if you need those.
## Quick start
### ring tunnel (no root, fastest)
```bash
cargo bench --bench tx_throughput
```
Smoke run (faster iteration):
```bash
TX_THROUGHPUT_MEASUREMENT_SECS=2 \
TX_THROUGHPUT_WARMUP_SECS=1 \
TX_THROUGHPUT_SAMPLE_SIZE=10 \
cargo bench --bench tx_throughput -- --quiet
```
### tcp / udp tunnels (requires Docker + root)
The benchmark creates a Docker network and registers each container's netns
under `/var/run/netns`, which requires root. Run the whole command under
`sudo`:
```bash
sudo TX_THROUGHPUT_TUNNEL=tcp \
TX_THROUGHPUT_MEASUREMENT_SECS=5 \
TX_THROUGHPUT_WARMUP_SECS=2 \
TX_THROUGHPUT_INFLIGHT=64 \
cargo bench --bench tx_throughput -- --quiet
sudo TX_THROUGHPUT_TUNNEL=udp cargo bench --bench tx_throughput -- --quiet
```
> If `sudo` cannot find `cargo`, use `sudo -E` or the absolute path
> (`$(which cargo)`).
## Environment variables
| Variable | Default | Notes |
| -------------------------------- | --------------------- | -------------------------------------- |
| `TX_THROUGHPUT_TUNNEL` | `ring` | `ring` / `tcp` / `udp` |
| `TX_THROUGHPUT_PKT_SIZE` | `1400` | IP total length in bytes |
| `TX_THROUGHPUT_WORKER_THREADS` | `4` | tokio worker threads |
| `TX_THROUGHPUT_INFLIGHT` | `64` | saturate-mode concurrency (task count) |
| `TX_THROUGHPUT_TUNNEL_PORT` | `35521` | tcp/udp listen port |
| `TX_THROUGHPUT_MEASUREMENT_SECS` | `10` | Criterion `measurement_time` |
| `TX_THROUGHPUT_WARMUP_SECS` | `3` | Criterion `warm_up_time` |
| `TX_THROUGHPUT_SAMPLE_SIZE` | `10` | Criterion `sample_size` (min 10) |
| `TX_THROUGHPUT_DOCKER_IMAGE` | `busybox:latest` | tcp/udp only |
| `TX_THROUGHPUT_DOCKER_NET` | `easytier-bench-<id>` | auto-generated unique name |
| `TX_THROUGHPUT_DOCKER_SUBNET` | `172.31.250.0/24` | |
| `TX_THROUGHPUT_DOCKER_IP_A` | `172.31.250.2` | |
| `TX_THROUGHPUT_DOCKER_IP_B` | `172.31.250.3` | |
## Parameter sweeps
```bash
# Packet size
for sz in 64 256 1400 9000; do
TX_THROUGHPUT_PKT_SIZE=$sz cargo bench --bench tx_throughput -- --quick
done
# Inflight depth (self-check: depth=1 should match serial baseline)
for d in 1 4 16 64 256; do
TX_THROUGHPUT_INFLIGHT=$d cargo bench --bench tx_throughput -- --quick
done
# Worker threads
for w in 1 2 4 8; do
TX_THROUGHPUT_WORKER_THREADS=$w cargo bench --bench tx_throughput -- --quick
done
```
## Interpreting results
- **`<tunnel>`** reports per-packet latency. Lower is better. Throughput
column here is "what one in-flight sender sustains".
- **`<tunnel>-saturate`** reports aggregate throughput across
`TX_THROUGHPUT_INFLIGHT` concurrent senders. If this matches the serial
baseline, the TX path is bottlenecked on an internal serialization point
(lock, single-threaded queue, etc.) rather than CPU or link bandwidth.
### Known finding (ring, single peer)
On the ring tunnel with a single destination peer, saturate does **not** beat
serial (observed ~277 MiB/s saturate vs ~288 MiB/s serial on a 4-worker
runtime). This points to a serialization point inside the peer-connection TX
path. Tunnels with real I/O await points (tcp/udp via Docker) are expected to
show a saturate > serial gap; verify with the sudo commands above.
## Output artifacts
Criterion writes HTML reports + SVG plots under
`easytier/target/criterion/`. Open `tx_throughput/<tunnel>/report/index.html`
or `.../<tunnel>-saturate/report/index.html` in a browser to inspect
distributions and regressions across runs.
@@ -1,65 +0,0 @@
use std::hint::black_box;
use std::time::Duration;
use criterion::{BatchSize, Criterion, Throughput, criterion_group, criterion_main};
use easytier::tunnel::packet_def::ZCPacket;
const PAYLOAD_SIZES: &[usize] = &[1280, 4096];
fn env_parse<T: std::str::FromStr>(key: &str, default: T) -> T {
std::env::var(key)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(default)
}
fn bench_payload_bytes(c: &mut Criterion) {
let mut group = c.benchmark_group("payload_bytes");
for &size in PAYLOAD_SIZES {
let data = vec![0u8; size];
group.throughput(Throughput::Bytes(size as u64));
group.bench_with_input(format!("{size}"), &data, |b, data| {
b.iter_batched(
|| ZCPacket::new_with_payload(black_box(data)),
|p| black_box(p).payload_bytes(),
BatchSize::SmallInput,
)
});
}
group.finish();
}
fn bench_tunnel_payload_bytes(c: &mut Criterion) {
let mut group = c.benchmark_group("tunnel_payload_bytes");
for &size in PAYLOAD_SIZES {
let data = vec![0u8; size];
group.throughput(Throughput::Bytes(size as u64));
group.bench_with_input(format!("{size}"), &data, |b, data| {
b.iter_batched(
|| ZCPacket::new_with_payload(black_box(data)),
|p| black_box(p).tunnel_payload_bytes(),
BatchSize::SmallInput,
)
});
}
group.finish();
}
fn criterion_config() -> Criterion {
let measurement_secs = env_parse("PACKET_BYTES_MEASUREMENT_SECS", 10u64);
let warmup_secs = env_parse("PACKET_BYTES_WARMUP_SECS", 3u64);
let sample_size = env_parse("PACKET_BYTES_SAMPLE_SIZE", 10usize).max(10);
Criterion::default()
.measurement_time(Duration::from_secs(measurement_secs))
.warm_up_time(Duration::from_secs(warmup_secs))
.sample_size(sample_size)
}
criterion_group! {
name = benches;
config = criterion_config();
targets = bench_payload_bytes, bench_tunnel_payload_bytes
}
criterion_main!(benches);
-472
View File
@@ -1,472 +0,0 @@
use std::{
net::IpAddr,
path::PathBuf,
process::{Command, Stdio},
str::FromStr,
sync::Arc,
sync::atomic::{AtomicU64, Ordering},
time::{Duration, Instant, SystemTime, UNIX_EPOCH},
};
use bytes::BytesMut;
use criterion::{Criterion, Throughput, criterion_group, criterion_main};
use easytier::{
common::config::{ConfigLoader, TomlConfigLoader},
instance::instance::Instance,
tunnel::{
packet_def::ZCPacket, ring::RingTunnelConnector, tcp::TcpTunnelConnector,
udp::UdpTunnelConnector,
},
};
const VIRTUAL_IP_A: &str = "10.144.144.1";
const VIRTUAL_IP_B: &str = "10.144.144.2";
const DEFAULT_DOCKER_SUBNET: &str = "172.31.250.0/24";
const DEFAULT_DOCKER_IP_A: &str = "172.31.250.2";
const DEFAULT_DOCKER_IP_B: &str = "172.31.250.3";
const DEFAULT_TUNNEL_PORT: u16 = 35521;
#[derive(Clone, Copy, Debug)]
enum TunnelKind {
Ring,
Tcp,
Udp,
}
impl TunnelKind {
fn as_str(self) -> &'static str {
match self {
TunnelKind::Ring => "ring",
TunnelKind::Tcp => "tcp",
TunnelKind::Udp => "udp",
}
}
}
impl FromStr for TunnelKind {
type Err = String;
fn from_str(value: &str) -> Result<Self, Self::Err> {
match value {
"ring" => Ok(TunnelKind::Ring),
"tcp" => Ok(TunnelKind::Tcp),
"udp" => Ok(TunnelKind::Udp),
other => Err(format!(
"unsupported TX_THROUGHPUT_TUNNEL={other:?}; expected ring, tcp, or udp"
)),
}
}
}
struct BenchTopology {
_docker: Option<DockerNetns>,
inst_a: Instance,
_inst_b: Instance,
dst: IpAddr,
packet: ZCPacket,
}
struct DockerNetns {
network: String,
container_a: String,
container_b: String,
netns_a: String,
netns_b: String,
ip_a: String,
netns_a_path: PathBuf,
netns_b_path: PathBuf,
}
impl DockerNetns {
fn create() -> Self {
let id = unique_id();
let image = env_string("TX_THROUGHPUT_DOCKER_IMAGE", "busybox:latest");
let network = env_string("TX_THROUGHPUT_DOCKER_NET", &format!("easytier-bench-{id}"));
let subnet = env_string("TX_THROUGHPUT_DOCKER_SUBNET", DEFAULT_DOCKER_SUBNET);
let ip_a = env_string("TX_THROUGHPUT_DOCKER_IP_A", DEFAULT_DOCKER_IP_A);
let ip_b = env_string("TX_THROUGHPUT_DOCKER_IP_B", DEFAULT_DOCKER_IP_B);
let container_a = format!("easytier-bench-a-{id}");
let container_b = format!("easytier-bench-b-{id}");
let netns_a = format!("easytier-bench-a-{id}");
let netns_b = format!("easytier-bench-b-{id}");
docker(&[
"network", "create", "--driver", "bridge", "--subnet", &subnet, &network,
]);
let mut docker_netns = Self {
network,
container_a,
container_b,
netns_a,
netns_b,
ip_a: ip_a.clone(),
netns_a_path: PathBuf::new(),
netns_b_path: PathBuf::new(),
};
docker_netns.start_container(&docker_netns.container_a, &ip_a, &image);
docker_netns.start_container(&docker_netns.container_b, &ip_b, &image);
let pid_a = docker(&["inspect", "-f", "{{.State.Pid}}", &docker_netns.container_a]);
let pid_b = docker(&["inspect", "-f", "{{.State.Pid}}", &docker_netns.container_b]);
docker_netns.netns_a_path = register_netns(&docker_netns.netns_a, &pid_a);
docker_netns.netns_b_path = register_netns(&docker_netns.netns_b, &pid_b);
docker_netns
}
fn start_container(&self, name: &str, ip: &str, image: &str) {
docker(&[
"run",
"-d",
"--name",
name,
"--network",
&self.network,
"--ip",
ip,
image,
"sleep",
"3600",
]);
}
}
impl Drop for DockerNetns {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.netns_a_path);
let _ = std::fs::remove_file(&self.netns_b_path);
docker_ignore(&["rm", "-f", &self.container_a, &self.container_b]);
docker_ignore(&["network", "rm", &self.network]);
}
}
fn bench_tx_throughput(c: &mut Criterion) {
let tunnel = env_string("TX_THROUGHPUT_TUNNEL", "ring")
.parse::<TunnelKind>()
.unwrap_or_else(|err| panic!("{err}"));
let packet_size = env_parse("TX_THROUGHPUT_PKT_SIZE", 1400usize);
const MIN_PKT_SIZE: usize = 28; // IPv4 (20) + UDP (8) header
assert!(
packet_size >= MIN_PKT_SIZE,
"TX_THROUGHPUT_PKT_SIZE={packet_size} is smaller than the minimum {MIN_PKT_SIZE} (IPv4+UDP headers)"
);
let worker_threads = env_parse("TX_THROUGHPUT_WORKER_THREADS", 4usize);
let inflight_depth = env_parse("TX_THROUGHPUT_INFLIGHT", 64usize).max(1);
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(worker_threads)
.enable_all()
.build()
.expect("create tokio runtime");
let topology = runtime.block_on(setup_topology(tunnel, packet_size));
let peer_manager = topology.inst_a.get_peer_manager();
let packet = topology.packet.clone();
let dst = topology.dst;
eprintln!(
"tx_throughput: tunnel={} inflight={} workers={} pkt_size={}",
tunnel.as_str(),
inflight_depth.max(1),
worker_threads,
packet_size
);
let mut group = c.benchmark_group("tx_throughput");
group.throughput(Throughput::Bytes(packet_size as u64));
// Serial baseline: one packet in flight at a time.
// Measures per-packet CPU cost (TX injection latency).
group.bench_function(tunnel.as_str(), |b| {
b.iter_custom(|iterations| {
let pm = peer_manager.clone();
let pkt = packet.clone();
runtime.block_on(async move {
let start = Instant::now();
for _ in 0..iterations {
pm.send_msg_by_ip(pkt.clone(), dst, false)
.await
.expect("send packet by EasyTier IP");
}
start.elapsed()
})
});
});
// Saturate: spawn TX_THROUGHPUT_INFLIGHT worker tasks, each independently
// pumping send_msg_by_ip. Work is distributed across tokio worker threads,
// exposing the peer manager + tunnel's true aggregate throughput ceiling.
// With TX_THROUGHPUT_INFLIGHT=1 it degrades to the serial baseline.
group.bench_function(format!("{}-saturate", tunnel.as_str()), |b| {
b.iter_custom(|iterations| {
let pm = peer_manager.clone();
let pkt = packet.clone();
let concurrency = inflight_depth.min(iterations as usize).max(1);
runtime.block_on(async move {
let counter = Arc::new(AtomicU64::new(iterations));
let start = Instant::now();
let mut handles = Vec::with_capacity(concurrency);
for _ in 0..concurrency {
let pm = pm.clone();
let pkt = pkt.clone();
let counter = counter.clone();
handles.push(tokio::spawn(async move {
loop {
if counter
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
if cur > 0 { Some(cur - 1) } else { None }
})
.is_err()
{
return;
}
pm.send_msg_by_ip(pkt.clone(), dst, false)
.await
.expect("send packet by EasyTier IP");
}
}));
}
for h in handles {
h.await.expect("saturate worker task panicked");
}
start.elapsed()
})
});
});
group.finish();
runtime.block_on(async move {
drop(topology);
});
}
async fn setup_topology(tunnel: TunnelKind, packet_size: usize) -> BenchTopology {
let tunnel_port = env_parse("TX_THROUGHPUT_TUNNEL_PORT", DEFAULT_TUNNEL_PORT);
let docker = match tunnel {
TunnelKind::Ring => None,
TunnelKind::Tcp | TunnelKind::Udp => Some(DockerNetns::create()),
};
let (netns_a, netns_b) = match &docker {
Some(docker) => (Some(docker.netns_a.clone()), Some(docker.netns_b.clone())),
None => (None, None),
};
let listeners_a = match tunnel {
TunnelKind::Ring => Vec::new(),
TunnelKind::Tcp | TunnelKind::Udp => vec![
format!("{}://0.0.0.0:{}", tunnel.as_str(), tunnel_port)
.parse()
.unwrap(),
],
};
let mut inst_a = Instance::new(no_tun_config("hot-a", VIRTUAL_IP_A, netns_a, listeners_a));
let mut inst_b = Instance::new(no_tun_config("hot-b", VIRTUAL_IP_B, netns_b, Vec::new()));
inst_a.run().await.expect("inst_a run");
inst_b.run().await.expect("inst_b run");
match tunnel {
TunnelKind::Ring => inst_b
.get_conn_manager()
.add_connector(RingTunnelConnector::new(
format!("ring://{}", inst_a.id()).parse().unwrap(),
)),
TunnelKind::Tcp => inst_b
.get_conn_manager()
.add_connector(TcpTunnelConnector::new(
format!(
"tcp://{}:{}",
docker.as_ref().expect("tcp benchmark needs Docker").ip_a,
tunnel_port
)
.parse()
.unwrap(),
)),
TunnelKind::Udp => inst_b
.get_conn_manager()
.add_connector(UdpTunnelConnector::new(
format!(
"udp://{}:{}",
docker.as_ref().expect("udp benchmark needs Docker").ip_a,
tunnel_port
)
.parse()
.unwrap(),
)),
}
wait_for_routes(&inst_a, &inst_b).await;
BenchTopology {
_docker: docker,
inst_a,
_inst_b: inst_b,
dst: VIRTUAL_IP_B.parse().unwrap(),
packet: make_data_packet(VIRTUAL_IP_A, VIRTUAL_IP_B, packet_size),
}
}
async fn wait_for_routes(inst_a: &Instance, inst_b: &Instance) {
tokio::time::timeout(Duration::from_secs(15), async {
loop {
let routes_a = inst_a.get_peer_manager().list_routes().await;
let routes_b = inst_b.get_peer_manager().list_routes().await;
if !routes_a.is_empty() && !routes_b.is_empty() {
return;
}
tokio::time::sleep(Duration::from_millis(500)).await;
}
})
.await
.expect("EasyTier routes did not converge within 15s");
}
fn make_data_packet(src: &str, dst: &str, total_size: usize) -> ZCPacket {
use std::net::Ipv4Addr;
let hdr_len = 28;
let payload_len = total_size.saturating_sub(hdr_len);
let ip_total_len = (hdr_len + payload_len) as u16;
let mut buf = BytesMut::with_capacity(total_size);
buf.extend_from_slice(&[
0x45,
0x00,
(ip_total_len >> 8) as u8,
(ip_total_len & 0xff) as u8,
0x00,
0x00,
0x40,
0x00,
0x40,
0x11,
0x00,
0x00,
]);
let src: Ipv4Addr = src.parse().unwrap();
buf.extend_from_slice(&src.octets());
let dst: Ipv4Addr = dst.parse().unwrap();
buf.extend_from_slice(&dst.octets());
let udp_len = (8 + payload_len) as u16;
buf.extend_from_slice(&[
0x30,
0x39,
0xd4,
0x31,
(udp_len >> 8) as u8,
(udp_len & 0xff) as u8,
0x00,
0x00,
]);
buf.resize(total_size, 0xaa);
ZCPacket::new_with_payload(&buf)
}
fn no_tun_config(
name: &str,
ipv4: &str,
netns: Option<String>,
listeners: Vec<url::Url>,
) -> TomlConfigLoader {
let config = TomlConfigLoader::default();
config.set_inst_name(name.to_owned());
config.set_netns(netns);
config.set_ipv4(Some(ipv4.parse().unwrap()));
config.set_listeners(listeners);
let mut flags = config.get_flags();
flags.no_tun = true;
config.set_flags(flags);
config
}
fn register_netns(name: &str, pid: &str) -> PathBuf {
#[cfg(target_os = "linux")]
{
let dir = PathBuf::from("/var/run/netns");
std::fs::create_dir_all(&dir).expect("create /var/run/netns");
let path = dir.join(name);
let _ = std::fs::remove_file(&path);
std::os::unix::fs::symlink(format!("/proc/{pid}/ns/net"), &path)
.expect("link Docker netns into /var/run/netns");
path
}
#[cfg(not(target_os = "linux"))]
{
let _ = (name, pid);
panic!("Docker netns benchmark requires Linux");
}
}
fn docker(args: &[&str]) -> String {
let output = Command::new("docker")
.args(args)
.output()
.unwrap_or_else(|err| panic!("failed to run docker {args:?}: {err}"));
if !output.status.success() {
panic!(
"docker {:?} failed with status {:?}: {}",
args,
output.status.code(),
String::from_utf8_lossy(&output.stderr)
);
}
String::from_utf8_lossy(&output.stdout).trim().to_owned()
}
fn docker_ignore(args: &[&str]) {
let _ = Command::new("docker")
.args(args)
.stdout(Stdio::null())
.stderr(Stdio::null())
.status();
}
fn env_string(name: &str, default: &str) -> String {
std::env::var(name).unwrap_or_else(|_| default.to_owned())
}
fn env_parse<T>(name: &str, default: T) -> T
where
T: FromStr,
T::Err: std::fmt::Display,
{
match std::env::var(name) {
Ok(value) => value
.parse()
.unwrap_or_else(|err| panic!("invalid {name}={value:?}: {err}")),
Err(_) => default,
}
}
fn criterion_config() -> Criterion {
let measurement_secs = env_parse("TX_THROUGHPUT_MEASUREMENT_SECS", 10u64);
let warmup_secs = env_parse("TX_THROUGHPUT_WARMUP_SECS", 3u64);
let sample_size = env_parse("TX_THROUGHPUT_SAMPLE_SIZE", 10usize).max(10);
Criterion::default()
.measurement_time(Duration::from_secs(measurement_secs))
.warm_up_time(Duration::from_secs(warmup_secs))
.sample_size(sample_size)
}
fn unique_id() -> String {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock before UNIX epoch")
.as_nanos();
format!("{}-{nanos}", std::process::id())
}
criterion_group! {
name = benches;
config = criterion_config();
targets = bench_tx_throughput
}
criterion_main!(benches);
+2 -2
View File
@@ -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: "最初要连接的对等节点"
+5 -7
View File
@@ -3,11 +3,9 @@ use std::{
net::{IpAddr, SocketAddr},
str::FromStr as _,
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
time::{Duration, Instant, SystemTime, UNIX_EPOCH},
};
use quanta::Instant;
use crate::common::{config::ConfigLoader, global_ctx::ArcGlobalCtx, token_bucket::TokenBucket};
use crate::proto::acl::*;
use anyhow::Context as _;
@@ -109,10 +107,10 @@ impl AclCacheKey {
// Cache entry with timestamp for LRU cleanup
#[derive(Debug, Clone)]
pub(crate) struct AclCacheEntry {
pub struct AclCacheEntry {
pub action: Action,
pub matched_rule: RuleId,
pub last_access: Instant,
pub last_access: std::time::Instant,
// New fields to track rule characteristics for proper cache behavior
pub conn_track_key: Option<String>,
pub rate_limit_keys: Vec<RateLimitKey>,
@@ -412,7 +410,7 @@ impl AclProcessor {
}
// Remove oldest entries (LRU cleanup)
let mut entries: Vec<(AclCacheKey, Instant)> = cache
let mut entries: Vec<(AclCacheKey, std::time::Instant)> = cache
.iter()
.map(|entry| (entry.key().clone(), entry.value().last_access))
.collect();
@@ -433,7 +431,7 @@ impl AclProcessor {
);
}
pub(crate) fn process_packet_with_cache_entry(
pub fn process_packet_with_cache_entry(
&self,
packet_info: &PacketInfo,
cache_entry: &AclCacheEntry,
+55 -252
View File
@@ -6,11 +6,9 @@ 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;
use prost_reflect::{DynamicMessage, ReflectMessage, SerializeOptions};
use serde::{Deserialize, Serialize};
use strum::{Display, EnumString, VariantArray};
use tokio::io::AsyncReadExt as _;
@@ -79,54 +77,6 @@ pub fn gen_default_flags() -> Flags {
}
}
fn flags_to_dynamic_message(flags: &Flags) -> DynamicMessage {
let mut message = DynamicMessage::new(flags.descriptor());
message
.transcode_from(flags)
.expect("FlagsInConfig should transcode to DynamicMessage");
message
}
fn flags_to_full_json_map(flags: &DynamicMessage) -> serde_json::Map<String, serde_json::Value> {
let options = SerializeOptions::new()
.use_proto_field_name(true)
.skip_default_fields(false);
match flags
.serialize_with_options(serde_json::value::Serializer, &options)
.expect("FlagsInConfig should serialize to JSON")
{
serde_json::Value::Object(map) => map,
_ => unreachable!("FlagsInConfig should serialize to a JSON object"),
}
}
fn flags_diff_from_default(flags: &Flags) -> serde_json::Map<String, serde_json::Value> {
let default_flags = gen_default_flags();
let default_message = flags_to_dynamic_message(&default_flags);
let current_message = flags_to_dynamic_message(flags);
let default_map = flags_to_full_json_map(&default_message);
let current_map = flags_to_full_json_map(&current_message);
current_message
.descriptor()
.fields()
.filter_map(|field| {
let key = field.name();
let value_changed = default_map.get(key) != current_map.get(key);
let presence_changed =
default_message.has_field(&field) != current_message.has_field(&field);
if value_changed || presence_changed {
current_map
.get(key)
.map(|value| (key.to_string(), value.clone()))
} else {
None
}
})
.collect()
}
fn mapped_listener_allows_implicit_port(url: &url::Url) -> bool {
TunnelScheme::try_from(url)
.ok()
@@ -235,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,
@@ -585,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>>,
@@ -619,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>>,
@@ -670,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 {
@@ -734,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()
}
}
@@ -861,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,
@@ -1142,9 +1059,28 @@ impl ConfigLoader for TomlConfigLoader {
}
fn dump(&self) -> String {
let default_flags_json = serde_json::to_string(&gen_default_flags()).unwrap();
let default_flags_hashmap =
serde_json::from_str::<serde_json::Map<String, serde_json::Value>>(&default_flags_json)
.unwrap();
let cur_flags_json = serde_json::to_string(&self.get_flags()).unwrap();
let cur_flags_hashmap =
serde_json::from_str::<serde_json::Map<String, serde_json::Value>>(&cur_flags_json)
.unwrap();
let mut flag_map: serde_json::Map<String, serde_json::Value> = Default::default();
for (key, value) in default_flags_hashmap {
if let Some(v) = cur_flags_hashmap.get(&key)
&& *v != value
{
flag_map.insert(key, v.clone());
}
}
let mut config = self.config.lock().unwrap().clone();
Self::normalize_config_source(&mut config);
config.flags = Some(flags_diff_from_default(&self.get_flags()));
config.flags = Some(flag_map);
if config.stun_servers == Some(StunInfoCollector::get_default_servers()) {
config.stun_servers = None;
}
@@ -1280,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)
@@ -1308,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;
@@ -1349,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.
@@ -1490,49 +1336,6 @@ socket_mark = 66
assert_eq!(cfg.get_flags().socket_mark, None);
}
#[test]
fn dump_preserves_flags_that_differ_from_easytier_defaults() {
let cfg = TomlConfigLoader::default();
let mut flags = gen_default_flags();
flags.dev_name = "et_test".to_string();
flags.enable_quic_proxy = true;
flags.disable_tcp_hole_punching = true;
flags.disable_sym_hole_punching = true;
flags.multi_thread = false;
flags.bind_device = false;
flags.enable_ipv6 = false;
flags.relay_network_whitelist = "".to_string();
flags.mtu = 0;
flags.socket_mark = Some(0);
cfg.set_flags(flags);
let dumped = cfg.dump();
assert!(dumped.contains("dev_name = \"et_test\""));
assert!(dumped.contains("enable_quic_proxy = true"));
assert!(dumped.contains("disable_tcp_hole_punching = true"));
assert!(dumped.contains("disable_sym_hole_punching = true"));
assert!(dumped.contains("multi_thread = false"));
assert!(dumped.contains("bind_device = false"));
assert!(dumped.contains("enable_ipv6 = false"));
assert!(dumped.contains("relay_network_whitelist = \"\""));
assert!(dumped.contains("mtu = 0"));
assert!(dumped.contains("socket_mark = 0"));
let reloaded = TomlConfigLoader::new_from_str(&dumped).unwrap();
let reloaded_flags = reloaded.get_flags();
assert_eq!(reloaded_flags.dev_name, "et_test");
assert!(reloaded_flags.enable_quic_proxy);
assert!(reloaded_flags.disable_tcp_hole_punching);
assert!(reloaded_flags.disable_sym_hole_punching);
assert!(!reloaded_flags.multi_thread);
assert!(!reloaded_flags.bind_device);
assert!(!reloaded_flags.enable_ipv6);
assert_eq!(reloaded_flags.relay_network_whitelist, "");
assert_eq!(reloaded_flags.mtu, 0);
assert_eq!(reloaded_flags.socket_mark, Some(0));
}
#[test]
fn test_stun_servers_config() {
let config = TomlConfigLoader::default();
-50
View File
@@ -219,7 +219,6 @@ pub struct GlobalCtx {
running_listeners: Mutex<Vec<url::Url>>,
advertised_ipv6_public_addr_prefix: Mutex<Option<cidr::Ipv6Cidr>>,
tun_device_name: Mutex<Option<String>>,
flags: ArcSwap<Flags>,
@@ -337,7 +336,6 @@ impl GlobalCtx {
running_listeners: Mutex::new(Vec::new()),
advertised_ipv6_public_addr_prefix: Mutex::new(None),
tun_device_name: Mutex::new(None),
flags: ArcSwap::new(Arc::new(flags)),
@@ -372,24 +370,6 @@ impl GlobalCtx {
}
}
fn set_tun_device_name(&self, name: Option<String>) {
*self.tun_device_name.lock().unwrap() = name;
}
pub(crate) fn set_tun_device_ready(&self, name: String) {
self.set_tun_device_name(Some(name.clone()));
self.issue_event(GlobalCtxEvent::TunDeviceReady(name));
}
pub(crate) fn set_tun_device_error(&self, error: String) {
self.set_tun_device_name(None);
self.issue_event(GlobalCtxEvent::TunDeviceError(error));
}
pub fn get_tun_device_name(&self) -> Option<String> {
self.tun_device_name.lock().unwrap().clone()
}
pub fn check_network_in_whitelist(&self, network_name: &str) -> Result<(), anyhow::Error> {
if self
.get_flags()
@@ -845,36 +825,6 @@ pub mod tests {
);
}
#[tokio::test]
async fn test_tun_device_name_tracks_explicit_runtime_state() {
let config = TomlConfigLoader::default();
let global_ctx = GlobalCtx::new(config);
assert_eq!(global_ctx.get_tun_device_name(), None);
global_ctx.issue_event(GlobalCtxEvent::TunDeviceReady("ignored".to_string()));
assert_eq!(global_ctx.get_tun_device_name(), None);
let mut subscriber = global_ctx.subscribe();
global_ctx.set_tun_device_ready("easytier0".to_string());
assert_eq!(
global_ctx.get_tun_device_name(),
Some("easytier0".to_string())
);
assert_eq!(
subscriber.recv().await.unwrap(),
GlobalCtxEvent::TunDeviceReady("easytier0".to_string())
);
global_ctx.set_tun_device_error("closed".to_string());
assert_eq!(global_ctx.get_tun_device_name(), None);
assert_eq!(
subscriber.recv().await.unwrap(),
GlobalCtxEvent::TunDeviceError("closed".to_string())
);
}
#[tokio::test]
async fn trusted_key_source_lookup_is_precise() {
let config = TomlConfigLoader::default();
-17
View File
@@ -177,20 +177,3 @@ pub(crate) fn list_ipv6_route_messages()
pub(crate) fn get_interface_index(name: &str) -> Result<u32, Error> {
netlink::NetlinkIfConfiger::get_interface_index(name)
}
#[cfg(target_os = "linux")]
pub(crate) fn add_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> {
netlink::NetlinkIfConfiger::add_ipv6_ndp_proxy(name, address)
}
#[cfg(target_os = "linux")]
pub(crate) fn remove_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> {
netlink::NetlinkIfConfiger::remove_ipv6_ndp_proxy(name, address)
}
#[cfg(target_os = "linux")]
pub(crate) fn list_ipv6_ndp_proxy(
name: &str,
) -> Result<std::collections::BTreeSet<Ipv6Addr>, Error> {
netlink::NetlinkIfConfiger::list_ipv6_ndp_proxy(name)
}
-104
View File
@@ -1,5 +1,4 @@
use std::{
collections::BTreeSet,
ffi::CString,
fmt::Debug,
net::{IpAddr, Ipv4Addr, Ipv6Addr},
@@ -17,10 +16,6 @@ use netlink_packet_core::{
use netlink_packet_route::{
AddressFamily, RouteNetlinkMessage,
address::{AddressAttribute, AddressMessage},
neighbour::{
NeighbourAddress, NeighbourAttribute, NeighbourFlags, NeighbourHeader, NeighbourMessage,
NeighbourState,
},
route::{
RouteAddress, RouteAttribute, RouteHeader, RouteMessage, RouteProtocol, RouteScope,
RouteType,
@@ -380,105 +375,6 @@ impl NetlinkIfConfiger {
pub(crate) fn list_ipv6_route_messages() -> Result<Vec<RouteMessage>, Error> {
Self::list_route_messages(AddressFamily::Inet6)
}
fn ipv6_ndp_proxy_message(name: &str, address: Ipv6Addr) -> Result<NeighbourMessage, Error> {
let mut message = NeighbourMessage::default();
message.header = NeighbourHeader {
family: AddressFamily::Inet6,
ifindex: Self::get_interface_index(name)?,
state: NeighbourState::Permanent,
flags: NeighbourFlags::Proxy,
kind: RouteType::Unicast,
};
message
.attributes
.push(NeighbourAttribute::Destination(NeighbourAddress::Inet6(
address,
)));
Ok(message)
}
pub(crate) fn add_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> {
send_netlink_req_and_wait_one_resp(
RouteNetlinkMessage::NewNeighbour(Self::ipv6_ndp_proxy_message(name, address)?),
false,
)
}
pub(crate) fn remove_ipv6_ndp_proxy(name: &str, address: Ipv6Addr) -> Result<(), Error> {
send_netlink_req_and_wait_one_resp(
RouteNetlinkMessage::DelNeighbour(Self::ipv6_ndp_proxy_message(name, address)?),
true,
)
}
fn list_neighbour_messages(
address_family: AddressFamily,
) -> Result<Vec<NeighbourMessage>, Error> {
let mut message = NeighbourMessage::default();
message.header.family = address_family;
message.header.flags = NeighbourFlags::Proxy;
let s = send_netlink_req(
RouteNetlinkMessage::GetNeighbour(message),
NLM_F_REQUEST | NLM_F_DUMP,
)?;
let mut ret_vec = vec![];
let mut resp = Vec::<u8>::new();
loop {
if resp.is_empty() {
let (new_resp, _) = s.recv_from_full()?;
resp = new_resp;
}
let ret = NetlinkMessage::<RouteNetlinkMessage>::deserialize(&resp)
.with_context(|| "Failed to deserialize netlink neighbour message")?;
resp = resp.split_off(ret.buffer_len());
tracing::debug!("net link response <<< {:?}", ret);
match ret.payload {
NetlinkPayload::Error(e) => {
if e.code == NonZero::new(0) {
continue;
} else {
return Err(e.to_io().into());
}
}
NetlinkPayload::InnerMessage(RouteNetlinkMessage::NewNeighbour(m)) => {
ret_vec.push(m);
}
NetlinkPayload::Done(_) => {
break;
}
p => {
tracing::error!("Unexpected netlink response: {:?}", p);
return Err(anyhow::anyhow!("Unexpected netlink response").into());
}
}
}
Ok(ret_vec)
}
pub(crate) fn list_ipv6_ndp_proxy(name: &str) -> Result<BTreeSet<Ipv6Addr>, Error> {
let ifindex = Self::get_interface_index(name)?;
Ok(Self::list_neighbour_messages(AddressFamily::Inet6)?
.into_iter()
.filter(|message| {
message.header.ifindex == ifindex
&& message.header.flags.contains(NeighbourFlags::Proxy)
})
.filter_map(|message| {
message.attributes.into_iter().find_map(|attr| match attr {
NeighbourAttribute::Destination(NeighbourAddress::Inet6(addr)) => Some(addr),
_ => None,
})
})
.collect())
}
}
#[async_trait]
+1 -2
View File
@@ -1,10 +1,9 @@
use dashmap::DashMap;
use quanta::Instant;
use serde::{Deserialize, Serialize};
use std::cell::UnsafeCell;
use std::fmt;
use std::sync::Arc;
use std::time::Duration;
use std::time::{Duration, Instant};
use tokio::time::interval;
use tokio_util::task::AbortOnDropHandle;
+2 -3
View File
@@ -2,13 +2,12 @@ use std::collections::BTreeSet;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::atomic::AtomicBool;
use std::sync::{Arc, RwLock};
use std::time::Duration;
use std::time::{Duration, Instant};
use crate::proto::common::{NatType, StunInfo};
use anyhow::Context;
use chrono::Local;
use crossbeam::atomic::AtomicCell;
use quanta::Instant;
use rand::seq::IteratorRandom;
use socket2::{SockAddr, SockRef};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@@ -1313,7 +1312,7 @@ impl StunInfoCollectorTrait for MockStunInfoCollector {
StunInfo {
udp_nat_type: self.udp_nat_type as i32,
tcp_nat_type: NatType::Unknown as i32,
last_update_time: Local::now().timestamp(),
last_update_time: std::time::Instant::now().elapsed().as_secs() as i64,
min_port: 100,
max_port: 200,
public_ip: vec!["127.0.0.1".to_string(), "::1".to_string()],
+21 -116
View File
@@ -8,11 +8,9 @@ use std::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
time::{Duration, Instant},
};
use quanta::Instant;
use crate::{
common::{
PeerId, dns::socket_addrs, error::Error, global_ctx::ArcGlobalCtx,
@@ -50,7 +48,6 @@ use url::Host;
pub const DIRECT_CONNECTOR_SERVICE_ID: u32 = 1;
pub const DIRECT_CONNECTOR_BLACKLIST_TIMEOUT_SEC: u64 = 300;
const MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16;
static TESTING: AtomicBool = AtomicBool::new(false);
@@ -85,56 +82,6 @@ fn is_usable_public_ipv6_candidate_with_mode(
&& !ip.is_multicast()))
}
fn push_ipv6_hole_punch_candidate(
candidates: &mut Vec<Ipv6Addr>,
ip: Ipv6Addr,
global_ctx: &ArcGlobalCtx,
limit: usize,
) {
if candidates.len() >= limit
|| !is_usable_public_ipv6_candidate(&ip, global_ctx)
|| candidates.contains(&ip)
{
return;
}
candidates.push(ip);
}
async fn collect_ipv6_hole_punch_candidates(global_ctx: &ArcGlobalCtx) -> Vec<Ipv6Addr> {
let mut candidates = Vec::new();
for ip in global_ctx
.get_stun_info_collector()
.get_stun_info()
.public_ip
.iter()
.filter_map(|ip| ip.parse::<Ipv6Addr>().ok())
{
push_ipv6_hole_punch_candidate(
&mut candidates,
ip,
global_ctx,
MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS,
);
}
let ip_list = global_ctx.get_ip_collector().collect_ip_addrs().await;
for ip in ip_list
.interface_ipv6s
.iter()
.chain(ip_list.public_ipv6.iter())
.map(|ip| Ipv6Addr::from(*ip))
{
push_ipv6_hole_punch_candidate(
&mut candidates,
ip,
global_ctx,
MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS,
);
}
candidates
}
#[async_trait::async_trait]
pub trait PeerManagerForDirectConnector {
async fn list_peers(&self) -> Vec<PeerId>;
@@ -204,8 +151,7 @@ impl DirectConnectorManagerData {
async fn remote_send_udp_hole_punch_packet(
&self,
dst_peer_id: PeerId,
connector_addrs: Vec<SocketAddr>,
preferred_src_ipv6: Option<Ipv6Addr>,
connector_addr: SocketAddr,
remote_url: &url::Url,
) -> Result<(), Error> {
if !matches_scheme!(remote_url, TunnelScheme::Ip(IpScheme::Udp)) {
@@ -236,17 +182,15 @@ impl DirectConnectorManagerData {
.send_udp_hole_punch_packet(
BaseController::default(),
SendUdpHolePunchPacketRequest {
connector_addr: connector_addrs.first().copied().map(Into::into),
listener_port: listener_port as u32,
preferred_src_ipv6: preferred_src_ipv6.map(Into::into),
connector_addrs: connector_addrs.into_iter().map(Into::into).collect(),
connector_addr: Some(connector_addr.into()),
},
)
.await
.with_context(|| {
format!(
"do rpc, send udp hole punch packet to peer {} at {} with preferred source {:?}",
dst_peer_id, remote_url, preferred_src_ipv6
"do rpc, send udp hole punch packet to peer {} at {}",
dst_peer_id, remote_url
)
})?;
@@ -263,41 +207,23 @@ impl DirectConnectorManagerData {
.await
.with_context(|| format!("failed to bind local socket for {}", remote_url))?,
);
let connector_ips = collect_ipv6_hole_punch_candidates(&self.global_ctx).await;
let connector_ip = self
.global_ctx
.get_stun_info_collector()
.get_stun_info()
.public_ip
.iter()
.filter_map(|ip| ip.parse::<Ipv6Addr>().ok())
.find(|ip| !self.global_ctx.is_ip_easytier_managed_ipv6(ip));
// ask remote to send v6 hole punch packet
// and no matter what the result is, continue to connect
if !connector_ips.is_empty() {
let local_port = local_socket.local_addr()?.port();
let connector_addrs = connector_ips
.into_iter()
.map(|ip| SocketAddr::new(IpAddr::V6(ip), local_port))
.collect::<Vec<_>>();
let preferred_src_ipv6 = match remote_url.host() {
Some(Host::Ipv6(ip)) => Some(ip),
_ => None,
};
tracing::debug!(
?connector_addrs,
?preferred_src_ipv6,
?remote_url,
"request remote IPv6 hole-punch packets"
);
if let Err(err) = self
.remote_send_udp_hole_punch_packet(
dst_peer_id,
connector_addrs,
preferred_src_ipv6,
remote_url,
)
.await
{
tracing::debug!(
?err,
?remote_url,
"remote IPv6 hole-punch packet request failed"
);
}
if let Some(connector_ip) = connector_ip {
let connector_addr =
SocketAddr::new(IpAddr::V6(connector_ip), local_socket.local_addr()?.port());
let _ = self
.remote_send_udp_hole_punch_packet(dst_peer_id, connector_addr, remote_url)
.await;
} else {
tracing::debug!(
?remote_url,
@@ -339,7 +265,7 @@ impl DirectConnectorManagerData {
.with_context(|| format!("failed to get udp port mapping for {}", remote_url))?;
let _ = self
.remote_send_udp_hole_punch_packet(dst_peer_id, vec![connector_addr], None, remote_url)
.remote_send_udp_hole_punch_packet(dst_peer_id, connector_addr, remote_url)
.await;
let udp_connector = UdpTunnelConnector::new(remote_url.clone());
@@ -890,7 +816,7 @@ mod tests {
tunnel::{IpScheme, TunnelScheme, matches_scheme},
};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use super::{TESTING, mapped_listener_port, resolve_mapped_listener_addrs};
@@ -912,27 +838,6 @@ mod tests {
));
}
#[tokio::test]
async fn ipv6_hole_punch_candidates_are_deduped_filtered_and_capped() {
let global_ctx = get_mock_global_ctx();
let managed_ipv6: cidr::Ipv6Inet = "2001:db8::2/128".parse().unwrap();
global_ctx.set_public_ipv6_routes(BTreeSet::from([managed_ipv6]));
let first: Ipv6Addr = "2001:db8::1".parse().unwrap();
let managed = managed_ipv6.address();
let second: Ipv6Addr = "2001:db8::3".parse().unwrap();
let third: Ipv6Addr = "2001:db8::4".parse().unwrap();
let mut candidates = Vec::new();
super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2);
super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2);
super::push_ipv6_hole_punch_candidate(&mut candidates, managed, &global_ctx, 2);
super::push_ipv6_hole_punch_candidate(&mut candidates, second, &global_ctx, 2);
super::push_ipv6_hole_punch_candidate(&mut candidates, third, &global_ctx, 2);
assert_eq!(candidates, vec![first, second]);
}
#[test]
fn udp_ipv6_url_matches_hole_punch_branch_condition() {
let remote_url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap();
+1 -2
View File
@@ -2,11 +2,10 @@ use std::{
collections::BTreeSet,
future::Future,
sync::{Arc, Weak},
time::Duration,
time::{Duration, Instant},
};
use dashmap::DashSet;
use quanta::Instant;
use tokio::{sync::mpsc, task::JoinSet, time::timeout};
use crate::{
+1 -2
View File
@@ -1,11 +1,10 @@
use std::{
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
sync::Arc,
time::Duration,
time::{Duration, Instant},
};
use anyhow::{Context, Error};
use quanta::Instant;
use rand::Rng as _;
use tokio::task::JoinSet;
@@ -1,11 +1,10 @@
use std::{
net::{IpAddr, SocketAddr, SocketAddrV4},
sync::Arc,
time::Duration,
time::{Duration, Instant},
};
use anyhow::Context;
use quanta::Instant;
use tokio::sync::Mutex;
use tokio_util::task::AbortOnDropHandle;
@@ -7,7 +7,6 @@ use std::{
use crossbeam::atomic::AtomicCell;
use dashmap::{DashMap, DashSet};
use guarden::defer;
use quanta::Instant;
use rand::seq::SliceRandom as _;
use tokio::{net::UdpSocket, sync::Mutex, task::JoinSet};
use tracing::{Instrument, Level, instrument};
@@ -357,9 +356,9 @@ pub(crate) struct UdpHolePunchListener {
_port_mapping_lease: Option<upnp::UdpPortMappingLease>,
conn_counter: Arc<Box<dyn TunnelConnCounter>>,
listen_time: Instant,
last_select_time: AtomicCell<Instant>,
last_active_time: Arc<AtomicCell<Instant>>,
listen_time: std::time::Instant,
last_select_time: AtomicCell<std::time::Instant>,
last_active_time: Arc<AtomicCell<std::time::Instant>>,
}
impl UdpHolePunchListener {
@@ -422,14 +421,14 @@ impl UdpHolePunchListener {
running_clone.store(false);
});
let last_active_time = Arc::new(AtomicCell::new(Instant::now()));
let last_active_time = Arc::new(AtomicCell::new(std::time::Instant::now()));
let conn_counter_clone = conn_counter.clone();
let last_active_time_clone = last_active_time.clone();
tasks.spawn(async move {
loop {
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
if conn_counter_clone.get().unwrap_or(0) != 0 {
last_active_time_clone.store(Instant::now());
last_active_time_clone.store(std::time::Instant::now());
}
}
});
@@ -445,14 +444,14 @@ impl UdpHolePunchListener {
_port_mapping_lease: port_mapping_lease,
conn_counter,
listen_time: Instant::now(),
last_select_time: AtomicCell::new(Instant::now()),
listen_time: std::time::Instant::now(),
last_select_time: AtomicCell::new(std::time::Instant::now()),
last_active_time,
})
}
pub async fn get_socket(&self) -> Arc<UdpSocket> {
self.last_select_time.store(Instant::now());
self.last_select_time.store(std::time::Instant::now());
self.socket.clone()
}
@@ -1,7 +1,9 @@
use std::{sync::Arc, time::Duration};
use std::{
sync::Arc,
time::{Duration, Instant},
};
use anyhow::Context;
use quanta::Instant;
use tokio::net::UdpSocket;
use tokio_util::task::AbortOnDropHandle;
+1 -2
View File
@@ -1,6 +1,6 @@
use std::{
sync::{Arc, atomic::AtomicBool},
time::Duration,
time::{Duration, Instant},
};
use anyhow::{Context, Error};
@@ -9,7 +9,6 @@ use common::{PunchHoleServerCommon, UdpNatType, UdpPunchClientMethod};
use cone::{PunchConeHoleClient, PunchConeHoleServer};
use dashmap::DashMap;
use once_cell::sync::Lazy;
use quanta::Instant;
use sym_to_cone::{PunchSymToConeHoleClient, PunchSymToConeHoleServer};
use tokio::{sync::Mutex, task::JoinHandle};
@@ -5,12 +5,11 @@ use std::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
time::{Duration, Instant},
};
use anyhow::Context;
use guarden::defer;
use quanta::Instant;
use rand::{Rng, seq::SliceRandom};
use tokio::{net::UdpSocket, sync::RwLock};
use tokio_util::task::AbortOnDropHandle;
+27 -11
View File
@@ -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))?;
};
}
+2 -3
View File
@@ -13,7 +13,6 @@ use pnet::packet::{
ip::IpNextHeaderProtocols,
ipv4::Ipv4Packet,
};
use quanta::Instant;
use socket2::Socket;
use tokio::{
sync::{Mutex, mpsc::UnboundedSender},
@@ -46,7 +45,7 @@ struct IcmpNatEntry {
src_peer_id: PeerId,
my_peer_id: PeerId,
src_ip: IpAddr,
start_time: Instant,
start_time: std::time::Instant,
mapped_dst_ip: std::net::Ipv4Addr,
}
@@ -61,7 +60,7 @@ impl IcmpNatEntry {
src_peer_id,
my_peer_id,
src_ip,
start_time: Instant::now(),
start_time: std::time::Instant::now(),
mapped_dst_ip,
})
}
+1 -2
View File
@@ -2,9 +2,8 @@ use dashmap::DashMap;
use pnet::packet::Packet;
use pnet::packet::ip::IpNextHeaderProtocol;
use pnet::packet::ipv4::{self, Ipv4Flags, Ipv4Packet, MutableIpv4Packet};
use quanta::Instant;
use std::net::Ipv4Addr;
use std::time::Duration;
use std::time::{Duration, Instant};
use crate::common::error::Error;
+4 -5
View File
@@ -1018,7 +1018,6 @@ impl TcpProxyRpc for QuicProxyDstRpcService {
mod tests {
use super::*;
use bytes::Buf;
use quanta::Instant;
/// Helper function: Create a pair of interconnected QuicSockets.
/// Data sent by socket_a will enter socket_b's rx, and vice versa.
@@ -1198,7 +1197,7 @@ mod tests {
// Accept unidirectional stream
let mut recv = connection.accept_uni().await.unwrap();
let start = Instant::now();
let start = std::time::Instant::now();
let mut received = 0;
// Loop read until the stream ends
@@ -1235,7 +1234,7 @@ mod tests {
let bytes_data = Bytes::from(data_chunk); // Use Bytes to avoid repeated allocation
println!("Client: Start sending {} MB...", TOTAL_SIZE / 1024 / 1024);
let start_send = Instant::now();
let start_send = std::time::Instant::now();
let chunks = TOTAL_SIZE / CHUNK_SIZE;
for _ in 0..chunks {
@@ -1277,7 +1276,7 @@ mod tests {
println!("Server: Accepted connection");
let mut stream_handles = Vec::new();
let start = Instant::now();
let start = std::time::Instant::now();
// Accept an expected number of streams
for i in 0..STREAM_COUNT {
@@ -1347,7 +1346,7 @@ mod tests {
STREAM_COUNT
);
let start_send = Instant::now();
let start_send = std::time::Instant::now();
let mut client_tasks = Vec::new();
// Start sending tasks concurrently
+57 -701
View File
@@ -5,13 +5,12 @@ use std::{
Arc, Weak,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
time::Duration,
time::{Duration, Instant},
};
use crossbeam::atomic::AtomicCell;
#[cfg(feature = "kcp")]
use kcp_sys::{endpoint::KcpEndpoint, stream::KcpStream};
use quanta::Instant;
use tokio_util::sync::{CancellationToken, DropGuard};
use tokio_util::task::AbortOnDropHandle;
@@ -32,7 +31,7 @@ use crate::{
tunnel::packet_def::{PacketType, ZCPacket},
};
use anyhow::Context;
use dashmap::{DashMap, mapref::entry::Entry};
use dashmap::DashMap;
use pnet::packet::{
Packet, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket,
};
@@ -164,87 +163,6 @@ struct Socks5Entry {
type Socks5EntrySet = Arc<DashMap<Socks5Entry, Socks5EntryData>>;
fn increment_entry_count(entry_count: &AtomicUsize) -> (usize, usize) {
let old_entry_count = entry_count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| {
count.checked_add(1)
})
.unwrap_or_else(|count| count);
(old_entry_count, old_entry_count.saturating_add(1))
}
fn decrement_entry_count(entry_count: &AtomicUsize) -> (usize, usize) {
decrement_entry_count_by(entry_count, 1)
}
fn decrement_entry_count_by(entry_count: &AtomicUsize, delta: usize) -> (usize, usize) {
if delta == 0 {
let current = entry_count.load(Ordering::Relaxed);
return (current, current);
}
let old_entry_count = entry_count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| {
Some(count.saturating_sub(delta))
})
.unwrap_or_else(|count| count);
(old_entry_count, old_entry_count.saturating_sub(delta))
}
fn insert_entry_and_increment_count(
entries: &Socks5EntrySet,
entry_count: &AtomicUsize,
entry: Socks5Entry,
data: Socks5EntryData,
) -> (bool, usize, usize) {
match entries.entry(entry) {
Entry::Occupied(mut occupied) => {
occupied.insert(data);
let current = entry_count.load(Ordering::Relaxed);
(true, current, current)
}
Entry::Vacant(vacant) => {
// Keep the count update inside the VacantEntry shard lock so bulk clear
// cannot observe the inserted entry before its count is reserved.
let (old_entry_count, new_entry_count) = increment_entry_count(entry_count);
vacant.insert(data);
(false, old_entry_count, new_entry_count)
}
}
}
fn try_insert_entry_and_increment_count(
entries: &Socks5EntrySet,
entry_count: &AtomicUsize,
entry: Socks5Entry,
data: Socks5EntryData,
) -> bool {
match entries.entry(entry) {
Entry::Occupied(_) => false,
Entry::Vacant(vacant) => {
// See insert_entry_and_increment_count for why the count is reserved first.
increment_entry_count(entry_count);
vacant.insert(data);
true
}
}
}
fn remove_entry_and_decrement_count(
entries: &Socks5EntrySet,
entry_count: &AtomicUsize,
entry: &Socks5Entry,
) -> (bool, usize, usize) {
let removed = entries.remove(entry).is_some();
let (old_entry_count, new_entry_count) = if removed {
decrement_entry_count(entry_count)
} else {
let current = entry_count.load(Ordering::Relaxed);
(current, current)
};
(removed, old_entry_count, new_entry_count)
}
struct SmolTcpConnector {
net: Arc<Net>,
entries: Socks5EntrySet,
@@ -271,20 +189,9 @@ impl AsyncTcpConnector for SmolTcpConnector {
entry_type: TCP_ENTRY,
};
*self.current_entry.lock().unwrap() = Some(entry.clone());
let (replaced, old_entry_count, new_entry_count) = insert_entry_and_increment_count(
&self.entries,
&self.entry_count,
entry.clone(),
Socks5EntryData::Tcp(tmp_listener),
);
tracing::trace!(
?entry,
replaced,
old_entry_count,
new_entry_count,
entries_len = self.entries.len(),
"socks5 inserted smoltcp tcp connector entry"
);
self.entries
.insert(entry, Socks5EntryData::Tcp(tmp_listener));
self.entry_count.fetch_add(1, Ordering::Relaxed);
if addr.ip() == local_addr {
let modified_addr =
@@ -312,16 +219,8 @@ impl Drop for SmolTcpConnector {
fn drop(&mut self) {
if let Some(entry) = self.current_entry.lock().unwrap().take() {
tracing::debug!("drop smoltcp connector entry {:?}", entry);
let (removed, old_entry_count, new_entry_count) =
remove_entry_and_decrement_count(&self.entries, &self.entry_count, &entry);
tracing::trace!(
?entry,
removed,
old_entry_count,
new_entry_count,
entries_len = self.entries.len(),
"socks5 removed smoltcp tcp connector entry"
);
self.entries.remove(&entry);
self.entry_count.fetch_sub(1, Ordering::Relaxed);
}
}
}
@@ -394,26 +293,11 @@ impl AsyncTcpConnector for Socks5AutoConnector {
addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), addr.port());
}
let has_smoltcp_net = self.smoltcp_net.is_some();
let dst_peers = if has_smoltcp_net && !addr.ip().is_loopback() {
Some(peer_mgr_arc.get_msg_dst_peer(&addr.ip()).await.0)
} else {
None
};
if !has_smoltcp_net
|| dst_peers.as_ref().is_some_and(Vec::is_empty)
if self.smoltcp_net.is_none()
|| peer_mgr_arc.get_msg_dst_peer(&addr.ip()).await.0.is_empty()
|| addr.ip().is_loopback()
{
// cannot find dst in virtual network, so try connect to dst directly
tracing::trace!(
?addr,
src_addr = ?self.src_addr,
has_smoltcp_net,
dst_peer_count = dst_peers.as_ref().map(Vec::len),
is_loopback = addr.ip().is_loopback(),
"socks5 auto connector falling back to kernel tcp connect"
);
return Ok(SocksTcpStream::Tcp(
tcp_connect_with_timeout(addr, timeout_s).await?,
));
@@ -425,51 +309,25 @@ impl AsyncTcpConnector for Socks5AutoConnector {
#[cfg(feature = "kcp")]
let connector: Box<dyn AsyncTcpConnector<S = SocksTcpStream> + Send> =
match (&self.kcp_endpoint, dst_allow_kcp) {
(Some(kcp_endpoint), true) => {
tracing::trace!(
?addr,
src_addr = ?self.src_addr,
dst_peer_count = dst_peers.as_ref().map(Vec::len),
"socks5 auto connector selected kcp"
);
Box::new(Socks5KcpConnector {
kcp_endpoint: kcp_endpoint.clone(),
peer_mgr: self.peer_mgr.clone(),
src_addr: self.src_addr,
})
}
(_, _) => {
tracing::trace!(
?addr,
src_addr = ?self.src_addr,
dst_peer_count = dst_peers.as_ref().map(Vec::len),
dst_allow_kcp,
has_kcp_endpoint = self.kcp_endpoint.is_some(),
"socks5 auto connector selected smoltcp"
);
Box::new(SmolTcpConnector {
net: self.smoltcp_net.clone().unwrap(),
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
current_entry: std::sync::Mutex::new(None),
})
}
(Some(kcp_endpoint), true) => Box::new(Socks5KcpConnector {
kcp_endpoint: kcp_endpoint.clone(),
peer_mgr: self.peer_mgr.clone(),
src_addr: self.src_addr,
}),
(_, _) => Box::new(SmolTcpConnector {
net: self.smoltcp_net.clone().unwrap(),
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
current_entry: std::sync::Mutex::new(None),
}),
};
#[cfg(not(feature = "kcp"))]
let connector = {
tracing::trace!(
?addr,
src_addr = ?self.src_addr,
dst_peer_count = dst_peers.as_ref().map(Vec::len),
"socks5 auto connector selected smoltcp"
);
Box::new(SmolTcpConnector {
net: self.smoltcp_net.clone().unwrap(),
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
current_entry: std::sync::Mutex::new(None),
})
};
let connector = Box::new(SmolTcpConnector {
net: self.smoltcp_net.clone().unwrap(),
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
current_entry: std::sync::Mutex::new(None),
});
let ret = connector.tcp_connect(addr, timeout_s).await;
self.inner_connector.lock().replace(Box::new(connector));
@@ -644,89 +502,25 @@ pub struct Socks5Server {
#[async_trait::async_trait]
impl PeerPacketFilter for Socks5Server {
async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option<ZCPacket> {
let entry_count = self.entry_count.load(Ordering::Relaxed);
let socks5_enabled = self.socks5_enabled.load(Ordering::Relaxed);
if entry_count == 0 && !socks5_enabled && self.entries.is_empty() {
if tracing::enabled!(tracing::Level::TRACE)
&& let Some(hdr) = packet.peer_manager_header()
&& matches!(
hdr.packet_type,
x if x == PacketType::Data as u8
|| x == PacketType::DataWithKcpSrcModified as u8
|| x == PacketType::DataWithQuicSrcModified as u8
)
{
if let Some(ipv4) = Ipv4Packet::new(packet.payload()) {
let (tcp_src_port, tcp_dst_port, tcp_flags) =
if ipv4.get_next_level_protocol() == IpNextHeaderProtocols::Tcp {
TcpPacket::new(ipv4.payload())
.map(|tcp| {
(
Some(tcp.get_source()),
Some(tcp.get_destination()),
Some(tcp.get_flags()),
)
})
.unwrap_or((None, None, None))
} else {
(None, None, None)
};
tracing::trace!(
packet_type = hdr.packet_type,
from_peer_id = hdr.from_peer_id.get(),
to_peer_id = hdr.to_peer_id.get(),
ipv4_src = %ipv4.get_source(),
ipv4_dst = %ipv4.get_destination(),
next_protocol = ?ipv4.get_next_level_protocol(),
?tcp_src_port,
?tcp_dst_port,
?tcp_flags,
entry_count,
socks5_enabled,
"socks5 fast gate passed packet from peer"
);
} else {
tracing::trace!(
packet_type = hdr.packet_type,
from_peer_id = hdr.from_peer_id.get(),
to_peer_id = hdr.to_peer_id.get(),
entry_count,
socks5_enabled,
"socks5 fast gate passed non-ipv4 packet from peer"
);
}
}
if self.entry_count.load(Ordering::Relaxed) == 0
&& !self.socks5_enabled.load(Ordering::Relaxed)
{
return Some(packet);
}
let hdr = packet.peer_manager_header().unwrap();
let is_modified_src_packet = matches!(
hdr.packet_type,
x if x == PacketType::DataWithKcpSrcModified as u8
|| x == PacketType::DataWithQuicSrcModified as u8
);
if hdr.packet_type != PacketType::Data as u8 && !is_modified_src_packet {
if hdr.packet_type != PacketType::Data as u8 {
return Some(packet);
}
if is_modified_src_packet && hdr.from_peer_id != hdr.to_peer_id {
tracing::trace!(
packet_type = hdr.packet_type,
from_peer_id = hdr.from_peer_id.get(),
to_peer_id = hdr.to_peer_id.get(),
"socks5 passed non-loopback modified-source packet from peer"
);
return Some(packet);
}
};
let payload_bytes = packet.payload();
let Some(ipv4) = Ipv4Packet::new(payload_bytes) else {
return Some(packet);
};
let ipv4 = Ipv4Packet::new(payload_bytes).unwrap();
if ipv4.get_version() != 4 {
return Some(packet);
}
let (entry_key, tcp_flags) = match ipv4.get_next_level_protocol() {
let entry_key = match ipv4.get_next_level_protocol() {
IpNextHeaderProtocols::Tcp => {
let Some(tcp_packet) = TcpPacket::new(ipv4.payload()) else {
return Some(packet);
@@ -751,11 +545,11 @@ impl PeerPacketFilter for Socks5Server {
entry_type: TCP_LISTEN_ENTRY,
}
};
(entry, Some(tcp_packet.get_flags()))
entry
}
IpNextHeaderProtocols::Udp => {
if IpReassembler::is_packet_fragmented(&ipv4) {
if IpReassembler::is_packet_fragmented(&ipv4) && !self.entries.is_empty() {
let ipv4_src: IpAddr = ipv4.get_source().into();
// only send to smoltcp if the ipv4 src is in the entries
let is_in_entries = self.entries.iter().any(|x| x.key().dst.ip() == ipv4_src);
@@ -767,19 +561,7 @@ impl PeerPacketFilter for Socks5Server {
if is_in_entries {
// if the packet is fragmented, no matther what the payload is, need send it to both smoltcp and kernel tun. because
// we cannot determine the udp port of the packet.
match self.packet_sender.try_send(packet.clone()) {
Ok(()) => tracing::trace!(
?ipv4_src,
entry_count = self.entry_count.load(Ordering::Relaxed),
"socks5 delivered fragmented packet from peer to smoltcp"
),
Err(err) => tracing::trace!(
?ipv4_src,
?err,
entry_count = self.entry_count.load(Ordering::Relaxed),
"socks5 failed to deliver fragmented packet from peer to smoltcp"
),
}
let _ = self.packet_sender.try_send(packet.clone()).ok();
}
return Some(packet);
}
@@ -787,17 +569,14 @@ impl PeerPacketFilter for Socks5Server {
let Some(udp_packet) = UdpPacket::new(ipv4.payload()) else {
return Some(packet);
};
(
Socks5Entry {
dst: SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
src: SocketAddr::new(
ipv4.get_destination().into(),
udp_packet.get_destination(),
),
entry_type: UDP_ENTRY,
},
None,
)
Socks5Entry {
dst: SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
src: SocketAddr::new(
ipv4.get_destination().into(),
udp_packet.get_destination(),
),
entry_type: UDP_ENTRY,
}
}
_ => {
return Some(packet);
@@ -805,41 +584,12 @@ impl PeerPacketFilter for Socks5Server {
};
if !self.entries.contains_key(&entry_key) {
tracing::trace!(
?entry_key,
?tcp_flags,
ipv4_src = %ipv4.get_source(),
ipv4_dst = %ipv4.get_destination(),
entry_count = self.entry_count.load(Ordering::Relaxed),
socks5_enabled = self.socks5_enabled.load(Ordering::Relaxed),
"socks5 no entry for packet from peer"
);
return Some(packet);
}
tracing::trace!(
?entry_key,
?tcp_flags,
?ipv4,
entry_count = self.entry_count.load(Ordering::Relaxed),
"socks5 found entry for packet from peer"
);
tracing::trace!(?entry_key, ?ipv4, "socks5 found entry for packet from peer");
match self.packet_sender.try_send(packet) {
Ok(()) => tracing::trace!(
?entry_key,
?tcp_flags,
entry_count = self.entry_count.load(Ordering::Relaxed),
"socks5 delivered packet from peer to smoltcp"
),
Err(err) => tracing::trace!(
?entry_key,
?tcp_flags,
?err,
entry_count = self.entry_count.load(Ordering::Relaxed),
"socks5 failed to deliver packet from peer to smoltcp"
),
}
let _ = self.packet_sender.try_send(packet).ok();
None
}
@@ -904,22 +654,11 @@ impl Socks5Server {
#[cfg(not(feature = "ffi-dataplane"))]
let data_plane_active = false;
let active_port_forwards = cancel_tokens.len();
let is_socks5_enabled = socks5_enabled.load(Ordering::Relaxed);
if active_port_forwards == 0 && !is_socks5_enabled && !data_plane_active {
let had_net = {
let mut net_guard = net.lock().await;
net_guard.take().is_some()
};
tracing::trace!(
had_net,
active_port_forwards,
is_socks5_enabled,
data_plane_active,
entry_count = entry_count.load(Ordering::Relaxed),
entries_len = entries.len(),
"socks5 net update waiting for consumers"
);
if cancel_tokens.is_empty()
&& !socks5_enabled.load(Ordering::Relaxed)
&& !data_plane_active
{
let _ = net.lock().await.take();
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(false);
port_forward_list_change_notifier.notified().await;
@@ -930,34 +669,13 @@ impl Socks5Server {
let cur_ipv4 = global_ctx.get_ipv4();
if prev_ipv4 != cur_ipv4 {
let old_ipv4 = prev_ipv4;
prev_ipv4 = cur_ipv4;
tracing::trace!(
?old_ipv4,
?cur_ipv4,
old_entry_count = entry_count.load(Ordering::Relaxed),
old_entries_len = entries.len(),
udp_client_count = udp_client_map.len(),
"socks5 net update resetting entries for ipv4 change"
);
let mut removed_entries = 0;
entries.retain(|_, _| {
removed_entries += 1;
entry_count.fetch_sub(1, Ordering::Relaxed);
false
});
let (_, new_entry_count) =
decrement_entry_count_by(&entry_count, removed_entries);
udp_client_map.clear();
tracing::trace!(
?old_ipv4,
?cur_ipv4,
removed_entries,
new_entry_count,
new_entries_len = entries.len(),
udp_client_count = udp_client_map.len(),
"socks5 net update reset entries complete"
);
if let Some(cur_ipv4) = cur_ipv4 {
net.lock().await.replace(Socks5ServerNet::new(
@@ -967,23 +685,12 @@ impl Socks5Server {
packet_recv.clone(),
entries.clone(),
));
tracing::trace!(
?cur_ipv4,
entry_count = entry_count.load(Ordering::Relaxed),
entries_len = entries.len(),
"socks5 net update installed smoltcp net"
);
// Wake any data-plane callers waiting in
// `wait_data_plane_net` for the smoltcp net to appear.
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(true);
} else {
let _ = net.lock().await.take();
tracing::trace!(
entry_count = entry_count.load(Ordering::Relaxed),
entries_len = entries.len(),
"socks5 net update removed smoltcp net"
);
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(false);
}
@@ -1066,13 +773,6 @@ impl Socks5Server {
peer_manager
.add_packet_process_pipeline(Box::new(self.clone()))
.await;
tracing::trace!(
cfg_count = cfgs.len(),
cancel_token_count = self.cancel_tokens.len(),
entry_count = self.entry_count.load(Ordering::Relaxed),
entries_len = self.entries.len(),
"socks5 peer packet pipeline registered"
);
self.run_net_update_task().await;
@@ -1105,7 +805,6 @@ impl Socks5Server {
connector: Box<dyn AsyncTcpConnector<S = SocksTcpStream> + Send>,
dst_addr: SocketAddr,
) {
tracing::trace!(?dst_addr, "port forward: connecting to destination");
let outgoing_socket = match connector.tcp_connect(dst_addr, 10).await {
Ok(socket) => socket,
Err(e) => {
@@ -1113,7 +812,6 @@ impl Socks5Server {
return;
}
};
tracing::trace!(?dst_addr, "port forward: connected to destination");
let mut outgoing_socket = outgoing_socket;
match tokio::io::copy_bidirectional(&mut incoming_socket, &mut outgoing_socket).await {
@@ -1200,30 +898,12 @@ impl Socks5Server {
dst_addr
);
let (smoltcp_net, net_ipv4) = {
let net_guard = net.lock().await;
(
net_guard.as_ref().map(|net| net.smoltcp_net.clone()),
net_guard.as_ref().map(|net| net.ipv4_addr),
)
};
tracing::trace!(
?bind_addr,
?dst_addr,
client_addr = ?addr,
has_smoltcp_net = smoltcp_net.is_some(),
?net_ipv4,
entry_count = entry_count.load(Ordering::Relaxed),
entries_len = entries.len(),
"port forward: preparing connector"
);
let connector = Socks5AutoConnector {
#[cfg(feature = "kcp")]
kcp_endpoint: kcp_endpoint.clone(),
peer_mgr: peer_mgr.clone(),
entries: entries.clone(),
smoltcp_net,
smoltcp_net: net.lock().await.as_ref().map(|net| net.smoltcp_net.clone()),
src_addr: addr,
entry_count: entry_count.clone(),
inner_connector: parking_lot::Mutex::new(None),
@@ -1359,12 +1039,11 @@ impl Socks5Server {
)
};
let socks_udp = Arc::new(sokcs_udp);
insert_entry_and_increment_count(
&entries,
&entry_count,
entries.insert(
client_info.entry_key.clone(),
Socks5EntryData::Udp((socks_udp.clone(), udp_client_key.clone())),
);
entry_count.fetch_add(1, Ordering::Relaxed);
let socks = socket.clone();
let client_addr = addr;
@@ -1427,18 +1106,16 @@ impl Socks5Server {
now.duration_since(client_info.last_active.load()).as_secs() < 600
});
udp_forward_task.retain(|k, _| udp_client_map.contains_key(k));
let mut removed_entries = 0;
entries.retain(|_, data| match data {
Socks5EntryData::Udp((_, udp_client_key)) => {
let keep = udp_client_map.contains_key(udp_client_key);
if !keep {
removed_entries += 1;
entry_count.fetch_sub(1, Ordering::Relaxed);
}
keep
}
_ => true,
});
decrement_entry_count_by(&entry_count, removed_entries);
udp_client_map.shrink_to_fit();
udp_forward_task.shrink_to_fit();
@@ -1450,324 +1127,3 @@ impl Socks5Server {
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use pnet::packet::{
MutablePacket,
ip::IpNextHeaderProtocols,
ipv4::{self, MutableIpv4Packet},
tcp::{self, MutableTcpPacket, TcpFlags},
};
use super::*;
use crate::peers::tests::create_mock_peer_manager;
fn build_tcp_packet(src: SocketAddr, dst: SocketAddr) -> Vec<u8> {
let mut buf = vec![0u8; 40];
let src_ip = match src.ip() {
IpAddr::V4(ip) => ip,
IpAddr::V6(_) => panic!("test only supports ipv4"),
};
let dst_ip = match dst.ip() {
IpAddr::V4(ip) => ip,
IpAddr::V6(_) => panic!("test only supports ipv4"),
};
{
let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap();
ip_packet.set_version(4);
ip_packet.set_header_length(5);
ip_packet.set_total_length(40);
ip_packet.set_ttl(64);
ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Tcp);
ip_packet.set_source(src_ip);
ip_packet.set_destination(dst_ip);
let mut tcp_packet = MutableTcpPacket::new(ip_packet.payload_mut()).unwrap();
tcp_packet.set_source(src.port());
tcp_packet.set_destination(dst.port());
tcp_packet.set_data_offset(5);
tcp_packet.set_flags(TcpFlags::SYN | TcpFlags::ACK);
tcp_packet.set_window(65535);
tcp_packet.set_checksum(tcp::ipv4_checksum(
&tcp_packet.to_immutable(),
&src_ip,
&dst_ip,
));
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
}
buf
}
fn build_udp_followup_fragment(src: Ipv4Addr, dst: Ipv4Addr) -> Vec<u8> {
let mut buf = vec![0u8; 28];
{
let mut ip_packet = MutableIpv4Packet::new(&mut buf).unwrap();
ip_packet.set_version(4);
ip_packet.set_header_length(5);
ip_packet.set_total_length(28);
ip_packet.set_ttl(64);
ip_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp);
ip_packet.set_fragment_offset(1);
ip_packet.set_source(src);
ip_packet.set_destination(dst);
ip_packet
.payload_mut()
.copy_from_slice(&[0xde, 0xad, 0xbe, 0xef, 0xca, 0xfe, 0xba, 0xbe]);
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
}
buf
}
#[tokio::test]
async fn socks5_consumes_modified_data_when_entry_matches() {
let peer_manager = create_mock_peer_manager().await;
let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None);
let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000);
let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22);
let entry = Socks5Entry {
src: local,
dst: remote,
entry_type: TCP_ENTRY,
};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
insert_entry_and_increment_count(
&server.entries,
&server.entry_count,
entry,
Socks5EntryData::Tcp(listener),
);
for packet_type in [
PacketType::DataWithKcpSrcModified,
PacketType::DataWithQuicSrcModified,
] {
let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local));
packet.fill_peer_manager_hdr(1, 1, packet_type as u8);
let result = server.try_process_packet_from_peer(packet).await;
assert!(result.is_none());
let mut receiver = server.packet_recv.lock().await;
let received = receiver.try_recv().unwrap();
assert_eq!(
received.peer_manager_header().unwrap().packet_type,
packet_type as u8
);
}
}
#[tokio::test]
async fn socks5_passes_through_unmatched_or_malformed_modified_data() {
let peer_manager = create_mock_peer_manager().await;
let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
insert_entry_and_increment_count(
&server.entries,
&server.entry_count,
Socks5Entry {
src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000),
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22),
entry_type: TCP_ENTRY,
},
Socks5EntryData::Tcp(listener),
);
let unmatched_local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40001);
let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22);
let mut unmatched_packet =
ZCPacket::new_with_payload(&build_tcp_packet(remote, unmatched_local));
unmatched_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8);
let result = server.try_process_packet_from_peer(unmatched_packet).await;
assert!(result.is_some());
let mut malformed_packet = ZCPacket::new_with_payload(&[0u8; 8]);
malformed_packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithQuicSrcModified as u8);
let result = server.try_process_packet_from_peer(malformed_packet).await;
assert!(result.is_some());
let mut receiver = server.packet_recv.lock().await;
assert!(receiver.try_recv().is_err());
}
#[tokio::test]
async fn socks5_passes_through_non_loopback_modified_data_even_when_entry_matches() {
let peer_manager = create_mock_peer_manager().await;
let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None);
let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000);
let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 22);
let entry = Socks5Entry {
src: local,
dst: remote,
entry_type: TCP_ENTRY,
};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
insert_entry_and_increment_count(
&server.entries,
&server.entry_count,
entry,
Socks5EntryData::Tcp(listener),
);
let mut packet = ZCPacket::new_with_payload(&build_tcp_packet(remote, local));
packet.fill_peer_manager_hdr(1, 2, PacketType::DataWithKcpSrcModified as u8);
let result = server.try_process_packet_from_peer(packet).await;
assert!(result.is_some());
let mut receiver = server.packet_recv.lock().await;
assert!(receiver.try_recv().is_err());
}
#[tokio::test]
async fn socks5_mirrors_fragmented_udp_even_when_entry_count_is_stale_zero() {
let peer_manager = create_mock_peer_manager().await;
let server = Socks5Server::new(peer_manager.get_global_ctx(), peer_manager, None);
let local = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 1)), 40000);
let remote = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 144, 144, 3)), 53);
let udp_socket = Arc::new(tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap());
server.entries.insert(
Socks5Entry {
src: local,
dst: remote,
entry_type: UDP_ENTRY,
},
Socks5EntryData::Udp((
Arc::new(SocksUdpSocket::UdpSocket(udp_socket)),
UdpClientKey {
client_addr: local,
dst_addr: remote,
},
)),
);
assert_eq!(server.entry_count.load(Ordering::Relaxed), 0);
let mut packet = ZCPacket::new_with_payload(&build_udp_followup_fragment(
match remote.ip() {
IpAddr::V4(ip) => ip,
IpAddr::V6(_) => unreachable!(),
},
match local.ip() {
IpAddr::V4(ip) => ip,
IpAddr::V6(_) => unreachable!(),
},
));
packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8);
let result = server.try_process_packet_from_peer(packet).await;
assert!(result.is_some());
let mut receiver = server.packet_recv.lock().await;
let received = receiver.try_recv().unwrap();
assert_eq!(
received.peer_manager_header().unwrap().packet_type,
PacketType::Data as u8
);
}
#[test]
fn decrement_entry_count_does_not_underflow() {
let entry_count = AtomicUsize::new(0);
let (old_entry_count, new_entry_count) = decrement_entry_count(&entry_count);
assert_eq!(old_entry_count, 0);
assert_eq!(new_entry_count, 0);
assert_eq!(entry_count.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn removing_missing_entry_does_not_decrement_entry_count() {
let entries = Arc::new(DashMap::new());
let entry_count = AtomicUsize::new(1);
let entry = Socks5Entry {
src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000),
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22),
entry_type: TCP_ENTRY,
};
let (removed, old_entry_count, new_entry_count) =
remove_entry_and_decrement_count(&entries, &entry_count, &entry);
assert!(!removed);
assert_eq!(old_entry_count, 1);
assert_eq!(new_entry_count, 1);
assert_eq!(entry_count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn removing_present_entry_decrements_entry_count_once() {
let entries = Arc::new(DashMap::new());
let entry_count = AtomicUsize::new(0);
let entry = Socks5Entry {
src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000),
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22),
entry_type: TCP_ENTRY,
};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
insert_entry_and_increment_count(
&entries,
&entry_count,
entry.clone(),
Socks5EntryData::Tcp(listener),
);
let (removed, old_entry_count, new_entry_count) =
remove_entry_and_decrement_count(&entries, &entry_count, &entry);
let (removed_again, old_entry_count_again, new_entry_count_again) =
remove_entry_and_decrement_count(&entries, &entry_count, &entry);
assert!(removed);
assert_eq!(old_entry_count, 1);
assert_eq!(new_entry_count, 0);
assert!(!removed_again);
assert_eq!(old_entry_count_again, 0);
assert_eq!(new_entry_count_again, 0);
assert_eq!(entry_count.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn replacing_present_entry_does_not_increment_entry_count() {
let entries = Arc::new(DashMap::new());
let entry_count = AtomicUsize::new(0);
let entry = Socks5Entry {
src: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 2)), 40000),
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 42, 0, 1)), 22),
entry_type: TCP_ENTRY,
};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let replacement = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let (replaced, old_entry_count, new_entry_count) = insert_entry_and_increment_count(
&entries,
&entry_count,
entry.clone(),
Socks5EntryData::Tcp(listener),
);
let (replaced_again, old_entry_count_again, new_entry_count_again) =
insert_entry_and_increment_count(
&entries,
&entry_count,
entry,
Socks5EntryData::Tcp(replacement),
);
assert!(!replaced);
assert_eq!(old_entry_count, 0);
assert_eq!(new_entry_count, 1);
assert!(replaced_again);
assert_eq!(old_entry_count_again, 1);
assert_eq!(new_entry_count_again, 1);
assert_eq!(entry_count.load(Ordering::Relaxed), 1);
}
}
+21 -27
View File
@@ -22,11 +22,11 @@ use std::{
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll},
time::Duration,
time::{Duration, Instant},
};
use anyhow::Context as _;
use quanta::Instant;
use dashmap::mapref::entry::Entry;
use tokio::io::{AsyncRead, AsyncWrite};
use crate::{common::error::Error, gateway::fast_socks5::server::AsyncTcpConnector};
@@ -34,7 +34,6 @@ use crate::{common::error::Error, gateway::fast_socks5::server::AsyncTcpConnecto
use super::{
Socks5AutoConnector, Socks5Entry, Socks5EntryData, Socks5EntrySet, Socks5Server,
SocksTcpStream, SocksUdpSocket, TCP_ENTRY, TCP_LISTEN_ENTRY, UDP_ENTRY, UdpClientKey,
decrement_entry_count, insert_entry_and_increment_count, try_insert_entry_and_increment_count,
};
use crate::gateway::tokio_smoltcp::{Net, TcpListener};
@@ -59,12 +58,12 @@ impl OwnedRouteEntry {
entry_count: Arc<AtomicUsize>,
entry: Socks5Entry,
) -> Self {
insert_entry_and_increment_count(
&entries,
&entry_count,
entry.clone(),
Socks5EntryData::DataPlaneRoute,
);
if entries
.insert(entry.clone(), Socks5EntryData::DataPlaneRoute)
.is_none()
{
entry_count.fetch_add(1, Ordering::Relaxed);
}
Self {
entries,
entry_count,
@@ -78,13 +77,12 @@ impl OwnedRouteEntry {
entry_count: Arc<AtomicUsize>,
entry: Socks5Entry,
) -> Option<Self> {
if !try_insert_entry_and_increment_count(
&entries,
&entry_count,
entry.clone(),
Socks5EntryData::DataPlaneRoute,
) {
return None;
match entries.entry(entry.clone()) {
Entry::Occupied(_) => return None,
Entry::Vacant(vacant) => {
vacant.insert(Socks5EntryData::DataPlaneRoute);
entry_count.fetch_add(1, Ordering::Relaxed);
}
}
Some(Self {
entries,
@@ -97,7 +95,7 @@ impl OwnedRouteEntry {
impl Drop for OwnedRouteEntry {
fn drop(&mut self) {
if self.entries.remove(&self.entry).is_some() {
decrement_entry_count(&self.entry_count);
self.entry_count.fetch_sub(1, Ordering::Relaxed);
}
}
}
@@ -225,18 +223,16 @@ impl DataPlaneUdpSocket {
dst: addr,
entry_type: UDP_ENTRY,
};
try_insert_entry_and_increment_count(
&self.entries,
&self.entry_count,
key,
Socks5EntryData::Udp((
if let Entry::Vacant(entry) = self.entries.entry(key) {
entry.insert(Socks5EntryData::Udp((
self.socket.clone(),
UdpClientKey {
client_addr: self.local_addr,
dst_addr: addr,
},
)),
);
)));
self.entry_count.fetch_add(1, Ordering::Relaxed);
}
self.socket.send_to(buf, addr).await
}
@@ -247,15 +243,13 @@ impl DataPlaneUdpSocket {
impl Drop for DataPlaneUdpSocket {
fn drop(&mut self) {
let mut removed_entries = 0;
self.entries.retain(|_, data| match data {
Socks5EntryData::Udp((socket, _)) if Arc::ptr_eq(socket, &self.socket) => {
removed_entries += 1;
self.entry_count.fetch_sub(1, Ordering::Relaxed);
false
}
_ => true,
});
super::decrement_entry_count_by(&self.entry_count, removed_entries);
}
}
+1 -2
View File
@@ -8,12 +8,11 @@ use pnet::packet::Packet;
use pnet::packet::ip::IpNextHeaderProtocols;
use pnet::packet::ipv4::{Ipv4Packet, MutableIpv4Packet};
use pnet::packet::tcp::{MutableTcpPacket, TcpPacket, ipv4_checksum};
use quanta::Instant;
use socket2::{SockRef, TcpKeepalive};
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::atomic::{AtomicBool, AtomicU16};
use std::sync::{Arc, Weak};
use std::time::Duration;
use std::time::{Duration, Instant};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, copy_bidirectional};
use tokio::net::{TcpListener, TcpSocket, TcpStream};
use tokio::sync::{Mutex, mpsc};
+5 -6
View File
@@ -14,7 +14,6 @@ use pnet::packet::{
ipv4::Ipv4Packet,
udp::{self, MutableUdpPacket},
};
use quanta::Instant;
use tokio::sync::mpsc::{Receiver, Sender, channel, error::TrySendError};
use tokio::{
net::UdpSocket,
@@ -61,8 +60,8 @@ struct UdpNatEntry {
socket: Option<UdpSocket>,
forward_task: Mutex<Option<JoinHandle<()>>>,
stopped: AtomicBool,
start_time: Instant,
last_active_time: AtomicCell<Instant>,
start_time: std::time::Instant,
last_active_time: AtomicCell<std::time::Instant>,
denied: bool,
}
@@ -86,8 +85,8 @@ impl UdpNatEntry {
socket,
forward_task: Mutex::new(None),
stopped: AtomicBool::new(false),
start_time: Instant::now(),
last_active_time: AtomicCell::new(Instant::now()),
start_time: std::time::Instant::now(),
last_active_time: AtomicCell::new(std::time::Instant::now()),
denied,
})
}
@@ -256,7 +255,7 @@ impl UdpNatEntry {
}
fn mark_active(&self) {
self.last_active_time.store(Instant::now());
self.last_active_time.store(std::time::Instant::now());
}
fn is_active(&self) -> bool {
+104 -84
View File
@@ -65,9 +65,9 @@ use crate::vpn_portal::{self, VpnPortal};
use super::dns_server::{MAGIC_DNS_FAKE_IP, runner::DnsRunner};
use super::listeners::ListenerManager;
use super::public_ipv6_provider::{
PublicIpv6ProviderReconcileTask, reconcile_public_ipv6_provider_runtime,
run_public_ipv6_provider_reconcile_task, should_run_public_ipv6_provider_reconcile,
validate_public_ipv6_config, validate_public_ipv6_config_values,
reconcile_public_ipv6_provider_runtime, run_public_ipv6_provider_reconcile_task,
should_run_public_ipv6_provider_reconcile, validate_public_ipv6_config,
validate_public_ipv6_config_values,
};
#[cfg(feature = "socks5")]
@@ -194,44 +194,6 @@ impl NicCtxContainer {
#[cfg(feature = "tun")]
type ArcNicCtx = Arc<Mutex<Option<NicCtxContainer>>>;
type ArcPublicIpv6ProviderTaskSlot = Arc<PublicIpv6ProviderTaskSlot>;
struct PublicIpv6ProviderTaskSlot {
task: Mutex<Option<PublicIpv6ProviderReconcileTask>>,
closing: AtomicBool,
}
impl PublicIpv6ProviderTaskSlot {
fn new() -> Self {
Self {
task: Mutex::new(None),
closing: AtomicBool::new(false),
}
}
async fn ensure_started(&self, global_ctx: &ArcGlobalCtx) {
let mut task = self.task.lock().await;
if self.closing.load(Ordering::Acquire) || task.is_some() {
return;
}
*task = run_public_ipv6_provider_reconcile_task(global_ctx);
}
async fn shutdown(&self) {
self.closing.store(true, Ordering::Release);
let task = self.task.lock().await.take();
if let Some(task) = task {
task.shutdown().await;
}
}
}
async fn ensure_public_ipv6_provider_reconcile_task(
global_ctx: &ArcGlobalCtx,
task_slot: &ArcPublicIpv6ProviderTaskSlot,
) {
task_slot.ensure_started(global_ctx).await;
}
pub struct InstanceRpcServerHook {
rpc_portal_whitelist: Vec<IpCidr>,
@@ -292,7 +254,6 @@ pub struct InstanceConfigPatcher {
socks5_server: Weak<Socks5Server>,
peer_manager: Weak<PeerManager>,
conn_manager: Weak<ManualConnectorManager>,
public_ipv6_provider_task: ArcPublicIpv6ProviderTaskSlot,
}
impl InstanceConfigPatcher {
@@ -363,6 +324,7 @@ impl InstanceConfigPatcher {
self.patch_mapped_listeners(patch.mapped_listeners).await?;
self.patch_connector(patch.connectors).await?;
let provider_reconcile_was_running = should_run_public_ipv6_provider_reconcile(&global_ctx);
let mut provider_config_changed = false;
if let Some(hostname) = patch.hostname {
global_ctx.set_hostname(hostname.clone());
@@ -400,12 +362,10 @@ impl InstanceConfigPatcher {
if provider_config_changed {
reconcile_public_ipv6_provider_runtime(&global_ctx).await;
if should_run_public_ipv6_provider_reconcile(&global_ctx) {
ensure_public_ipv6_provider_reconcile_task(
&global_ctx,
&self.public_ipv6_provider_task,
)
.await;
let provider_reconcile_should_run =
should_run_public_ipv6_provider_reconcile(&global_ctx);
if !provider_reconcile_was_running && provider_reconcile_should_run {
run_public_ipv6_provider_reconcile_task(&global_ctx);
}
}
@@ -523,22 +483,18 @@ impl InstanceConfigPatcher {
}
let global_ctx = weak_upgrade(&self.global_ctx)?;
for proxy_network_patch in proxy_networks {
let Some(cidr) = proxy_network_patch.cidr.map(|c| c.into()) else {
tracing::warn!("Proxy network cidr is None, skipping.");
continue;
};
let mapped_cidr: Option<cidr::Ipv4Cidr> =
proxy_network_patch.mapped_cidr.map(|s| s.into());
match ConfigPatchAction::try_from(proxy_network_patch.action) {
Ok(ConfigPatchAction::Add) => {
let Some(cidr) = proxy_network_patch.cidr.map(|c| c.into()) else {
tracing::warn!("Proxy network cidr is None, skipping add.");
continue;
};
let mapped_cidr: Option<cidr::Ipv4Cidr> =
proxy_network_patch.mapped_cidr.map(|s| s.into());
tracing::info!("Proxy network added: {}", cidr);
global_ctx.config.add_proxy_cidr(cidr, mapped_cidr)?;
}
Ok(ConfigPatchAction::Remove) => {
let Some(cidr) = proxy_network_patch.cidr.map(|c| c.into()) else {
tracing::warn!("Proxy network cidr is None, skipping remove.");
continue;
};
tracing::info!("Proxy network removed: {}", cidr);
global_ctx.config.remove_proxy_cidr(cidr);
}
@@ -687,7 +643,6 @@ pub struct Instance {
socks5_server: Arc<Socks5Server>,
proxy_cidrs_monitor: Option<AbortOnDropHandle<()>>,
public_ipv6_provider_task: ArcPublicIpv6ProviderTaskSlot,
global_ctx: ArcGlobalCtx,
}
@@ -775,7 +730,6 @@ impl Instance {
socks5_server,
proxy_cidrs_monitor: None,
public_ipv6_provider_task: Arc::new(PublicIpv6ProviderTaskSlot::new()),
global_ctx,
}
@@ -875,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());
@@ -910,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()
@@ -1076,11 +1038,7 @@ impl Instance {
.await?;
self.listener_manager.lock().await.run().await?;
self.peer_manager.run().await?;
ensure_public_ipv6_provider_reconcile_task(
&self.global_ctx,
&self.public_ipv6_provider_task,
)
.await;
run_public_ipv6_provider_reconcile_task(&self.global_ctx);
#[cfg(feature = "tun")]
{
@@ -1393,7 +1351,6 @@ impl Instance {
socks5_server: Arc::downgrade(&self.socks5_server),
peer_manager: Arc::downgrade(&self.peer_manager),
conn_manager: Arc::downgrade(&self.conn_manager),
public_ipv6_provider_task: self.public_ipv6_provider_task.clone(),
}
}
@@ -1649,7 +1606,6 @@ impl Instance {
}
pub async fn clear_resources(&mut self) {
self.public_ipv6_provider_task.shutdown().await;
self.peer_manager.clear_resources().await;
#[cfg(feature = "tun")]
let _ = self.nic_ctx.lock().await.take();
@@ -1835,21 +1791,6 @@ mod tests {
);
}
#[tokio::test]
async fn public_ipv6_provider_task_slot_does_not_restart_after_shutdown() {
let global_ctx = get_mock_global_ctx();
let slot = std::sync::Arc::new(super::PublicIpv6ProviderTaskSlot::new());
global_ctx.config.set_ipv6_public_addr_provider(true);
global_ctx
.config
.set_ipv6_public_addr_prefix(Some("2001:db8::/48".parse().unwrap()));
slot.shutdown().await;
super::ensure_public_ipv6_provider_reconcile_task(&global_ctx, &slot).await;
assert!(slot.task.lock().await.is_none());
}
#[tokio::test]
async fn validate_public_ipv6_patch_allows_enabling_auto_with_manual_ipv6() {
let global_ctx = get_mock_global_ctx();
@@ -1877,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()
);
}
}
+1 -1
View File
@@ -1,9 +1,9 @@
use std::collections::BTreeSet;
use std::sync::{Arc, Weak};
use std::time::Instant;
use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent};
use crate::peers::peer_manager::PeerManager;
use quanta::Instant;
use tokio_util::task::AbortOnDropHandle;
/// ProxyCidrsMonitor monitors changes in proxy CIDRs from peer routes
File diff suppressed because it is too large Load Diff
+9 -8
View File
@@ -24,7 +24,7 @@ use crate::{
};
use byteorder::WriteBytesExt as _;
use bytes::{Buf, BufMut, BytesMut};
use bytes::{BufMut, BytesMut};
use cidr::{Ipv4Inet, Ipv6Inet};
use futures::{SinkExt, Stream, StreamExt, lock::BiLock, ready};
use pin_project_lite::pin_project;
@@ -180,13 +180,12 @@ impl ZCPacketToBytes for TunZCPacketToBytes {
assert!(payload_offset >= 4);
let ret = if self.has_packet_info {
inner.advance(payload_offset - 4);
let mut inner = inner.split_off(payload_offset - 4);
let proto = infer_proto(&inner[4..]);
self.fill_packet_info(&mut inner[0..4], proto)?;
inner
} else {
inner.advance(payload_offset);
inner
inner.split_off(payload_offset)
};
tracing::debug!(?ret, ?payload_offset, "convert zc packet to tun packet");
@@ -1362,11 +1361,12 @@ impl NicCtx {
}
self.global_ctx
.set_tun_device_ready(nic.ifname().to_string());
.issue_event(GlobalCtxEvent::TunDeviceReady(nic.ifname().to_string()));
ret
}
Err(err) => {
self.global_ctx.set_tun_device_error(err.to_string());
self.global_ctx
.issue_event(GlobalCtxEvent::TunDeviceError(err.to_string()));
return Err(err);
}
}
@@ -1405,11 +1405,12 @@ impl NicCtx {
match nic.create_dev_for_mobile(tun_fd).await {
Ok(ret) => {
self.global_ctx
.set_tun_device_ready(nic.ifname().to_string());
.issue_event(GlobalCtxEvent::TunDeviceReady(nic.ifname().to_string()));
ret
}
Err(err) => {
self.global_ctx.set_tun_device_error(err.to_string());
self.global_ctx
.issue_event(GlobalCtxEvent::TunDeviceError(err.to_string()));
return Err(err);
}
}
+34 -183
View File
@@ -173,12 +173,7 @@ impl EasyTierLauncher {
#[cfg(mobile)]
Self::run_routine_for_mobile(&instance, &data, &mut tasks).await;
if let Err(err) = instance.run().await {
tasks.abort_all();
drop(tasks);
instance.clear_resources().await;
return Err(err.into());
}
instance.run().await?;
#[cfg(feature = "ffi-dataplane")]
data.data_plane
@@ -631,47 +626,6 @@ pub type NetworkingMethod = crate::proto::api::manage::NetworkingMethod;
pub type NetworkConfig = crate::proto::api::manage::NetworkConfig;
impl NetworkConfig {
fn parse_peer(peer: &manage::NetworkPeerConfig) -> Result<Option<PeerConfig>, anyhow::Error> {
let uri = peer.uri.trim();
if uri.is_empty() {
return Ok(None);
}
Ok(Some(PeerConfig {
uri: uri
.parse()
.with_context(|| format!("failed to parse peer uri: {}", uri))?,
peer_public_key: peer.peer_public_key.clone(),
}))
}
fn parse_peers(peers: &[manage::NetworkPeerConfig]) -> Result<Vec<PeerConfig>, anyhow::Error> {
let mut ret = Vec::new();
for peer in peers {
if let Some(peer) = Self::parse_peer(peer)? {
ret.push(peer);
}
}
Ok(ret)
}
fn parse_peer_urls(peer_urls: &[String]) -> Result<Vec<PeerConfig>, anyhow::Error> {
let mut peers = vec![];
for peer_url in peer_urls.iter() {
let peer_url = peer_url.trim();
if peer_url.is_empty() {
continue;
}
peers.push(PeerConfig {
uri: peer_url
.parse()
.with_context(|| format!("failed to parse peer uri: {}", peer_url))?,
peer_public_key: None,
});
}
Ok(peers)
}
pub fn gen_config(&self) -> Result<TomlConfigLoader, anyhow::Error> {
let cfg = TomlConfigLoader::default();
cfg.set_id(
@@ -683,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
@@ -727,23 +694,26 @@ impl NetworkConfig {
.unwrap_or_default()
{
NetworkingMethod::PublicServer => {
let peers = Self::parse_peers(&self.peers)?;
if peers.is_empty() {
let public_server_url = self.public_server_url.clone().unwrap_or_default();
cfg.set_peers(vec![PeerConfig {
uri: public_server_url.parse().with_context(|| {
format!("failed to parse public server uri: {}", public_server_url)
})?,
peer_public_key: None,
}]);
} else {
cfg.set_peers(peers);
}
let public_server_url = self.public_server_url.clone().unwrap_or_default();
cfg.set_peers(vec![PeerConfig {
uri: public_server_url.parse().with_context(|| {
format!("failed to parse public server uri: {}", public_server_url)
})?,
peer_public_key: None,
}]);
}
NetworkingMethod::Manual => {
let mut peers = Self::parse_peers(&self.peers)?;
if peers.is_empty() {
peers = Self::parse_peer_urls(&self.peer_urls)?;
let mut peers = vec![];
for peer_url in self.peer_urls.iter() {
if peer_url.is_empty() {
continue;
}
peers.push(PeerConfig {
uri: peer_url
.parse()
.with_context(|| format!("failed to parse peer uri: {}", peer_url))?,
peer_public_key: None,
});
}
if !peers.is_empty() {
cfg.set_peers(peers);
@@ -1062,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());
@@ -1087,13 +1058,6 @@ impl NetworkConfig {
result.networking_method = Some(NetworkingMethod::Manual as i32);
if !peers.is_empty() {
result.peer_urls = peers.iter().map(|p| p.uri.to_string()).collect();
result.peers = peers
.iter()
.map(|p| manage::NetworkPeerConfig {
uri: p.uri.to_string(),
peer_public_key: p.peer_public_key.clone(),
})
.collect();
}
result.listener_urls = config
@@ -1166,7 +1130,6 @@ impl NetworkConfig {
.get_credential_file()
.map(|path| path.to_string_lossy().into_owned());
let flags = config.get_flags();
let default_flags = default_config.get_flags();
result.latency_first = Some(flags.latency_first);
result.dev_name = Some(flags.dev_name.clone());
result.use_smoltcp = Some(flags.use_smoltcp);
@@ -1195,11 +1158,6 @@ impl NetworkConfig {
result.disable_sym_hole_punching = Some(flags.disable_sym_hole_punching);
result.enable_magic_dns = Some(flags.accept_dns);
result.mtu = Some(flags.mtu as i32);
result.data_compress_algo = (flags.data_compress_algo != default_flags.data_compress_algo)
.then_some(flags.data_compress_algo);
result.encryption_algorithm = (flags.encryption_algorithm
!= default_flags.encryption_algorithm)
.then_some(flags.encryption_algorithm.clone());
result.instance_recv_bps_limit =
(flags.instance_recv_bps_limit != u64::MAX).then_some(flags.instance_recv_bps_limit);
result.enable_private_mode = Some(flags.private_mode);
@@ -1229,7 +1187,7 @@ impl NetworkConfig {
mod tests {
use crate::{
common::config::{ConfigLoader, process_secure_mode_cfg},
proto::common::{CompressionAlgoPb, SecureModeConfig},
proto::common::SecureModeConfig,
};
use base64::prelude::{BASE64_STANDARD, Engine as _};
use rand::Rng;
@@ -1266,80 +1224,6 @@ mod tests {
Ok(())
}
#[test]
fn network_config_dump_preserves_web_flags() -> Result<(), anyhow::Error> {
let network_config = super::NetworkConfig {
instance_id: Some(uuid::Uuid::new_v4().to_string()),
dhcp: Some(true),
network_name: Some("demo".to_string()),
network_secret: Some("secret".to_string()),
networking_method: Some(crate::proto::api::manage::NetworkingMethod::Manual as i32),
peer_urls: vec!["tcp://1.2.3.4:11010".to_string()],
listener_urls: vec!["tcp://0.0.0.0:11010".to_string()],
dev_name: Some("et_test".to_string()),
enable_quic_proxy: Some(true),
disable_tcp_hole_punching: Some(true),
disable_sym_hole_punching: Some(true),
..Default::default()
};
let dumped = network_config.gen_config()?.dump();
assert!(dumped.contains("dev_name = \"et_test\""));
assert!(dumped.contains("enable_quic_proxy = true"));
assert!(dumped.contains("disable_tcp_hole_punching = true"));
assert!(dumped.contains("disable_sym_hole_punching = true"));
Ok(())
}
#[test]
fn test_network_config_conversion_preserves_peer_public_key() -> Result<(), anyhow::Error> {
let peer_url = "tcp://1.2.3.4:11010";
let peer_public_key = BASE64_STANDARD.encode([9u8; 32]);
let config = gen_default_config();
config.set_peers(vec![crate::common::config::PeerConfig {
uri: peer_url.parse()?,
peer_public_key: Some(peer_public_key.clone()),
}]);
let network_config = super::NetworkConfig::new_from_config(&config)?;
assert_eq!(network_config.peer_urls, vec![peer_url.to_string()]);
assert_eq!(network_config.peers.len(), 1);
assert_eq!(network_config.peers[0].uri, peer_url);
assert_eq!(
network_config.peers[0].peer_public_key.as_deref(),
Some(peer_public_key.as_str())
);
let generated_config = network_config.gen_config()?;
assert_eq!(generated_config.get_peers(), config.get_peers());
Ok(())
}
#[test]
fn network_config_gen_config_trims_legacy_peer_urls() -> Result<(), anyhow::Error> {
let network_config = super::NetworkConfig {
instance_id: Some(uuid::Uuid::new_v4().to_string()),
dhcp: Some(true),
networking_method: Some(crate::proto::api::manage::NetworkingMethod::Manual as i32),
peer_urls: vec![
" tcp://1.2.3.4:11010 ".to_string(),
" ".to_string(),
"\tudp://5.6.7.8:11010\n".to_string(),
],
..Default::default()
};
let generated_config = network_config.gen_config()?;
let peers = generated_config.get_peers();
assert_eq!(peers.len(), 2);
assert_eq!(peers[0].uri.as_str(), "tcp://1.2.3.4:11010");
assert_eq!(peers[1].uri.as_str(), "udp://5.6.7.8:11010");
Ok(())
}
#[test]
fn test_network_config_conversion_random() -> Result<(), anyhow::Error> {
let mut rng = rand::thread_rng();
@@ -1638,37 +1522,4 @@ mod tests {
Ok(())
}
#[test]
fn test_network_config_conversion_preserves_runtime_algorithm_flags()
-> Result<(), anyhow::Error> {
let config = gen_default_config();
let mut flags = config.get_flags();
flags.data_compress_algo = CompressionAlgoPb::Zstd.into();
flags.encryption_algorithm = "managed-test-algo".to_string();
config.set_flags(flags.clone());
let network_config = super::NetworkConfig::new_from_config(&config)?;
assert_eq!(
network_config.data_compress_algo,
Some(CompressionAlgoPb::Zstd as i32)
);
assert_eq!(
network_config.encryption_algorithm.as_deref(),
Some("managed-test-algo")
);
let generated_config = network_config.gen_config()?;
assert_eq!(
generated_config.get_flags().data_compress_algo,
flags.data_compress_algo
);
assert_eq!(
generated_config.get_flags().encryption_algorithm,
flags.encryption_algorithm
);
Ok(())
}
}
-5
View File
@@ -5,11 +5,6 @@ use std::io;
use clap::Command;
use clap_complete::{Generator, Shell};
// Re-export `Instant` at the crate root so public APIs that expose it
// (e.g. `Route::get_peer_info_last_update_time`) reference a deliberate
// public type rather than leaking an inaccessible one.
pub use quanta::Instant;
mod arch;
mod gateway;
pub mod instance;
+2 -3
View File
@@ -1,5 +1,6 @@
use std::net::{Ipv4Addr, Ipv6Addr};
use std::sync::atomic::Ordering;
use std::time::Instant;
use std::{
net::IpAddr,
sync::{Arc, atomic::AtomicBool},
@@ -11,7 +12,6 @@ use pnet::packet::ipv6::Ipv6Packet;
use pnet::packet::{
Packet as _, ip::IpNextHeaderProtocols, ipv4::Ipv4Packet, tcp::TcpPacket, udp::UdpPacket,
};
use quanta::Instant;
use crate::proto::acl::{AclStats, Protocol};
use crate::tunnel::packet_def::PacketType;
@@ -402,10 +402,9 @@ mod tests {
use std::{
net::{IpAddr, Ipv4Addr, Ipv6Addr},
sync::Arc,
time::Instant,
};
use quanta::Instant;
use crate::{
common::acl_processor::PacketInfo,
proto::acl::{Acl, ChainType, Protocol},
@@ -14,7 +14,7 @@ use std::{
};
use dashmap::{DashMap, DashSet};
use guarden::{Guard, defer};
use guarden::defer;
use tokio::{
sync::{
Mutex,
+15 -56
View File
@@ -1,4 +1,3 @@
use arc_swap::ArcSwapOption;
use crossbeam::atomic::AtomicCell;
use futures::{StreamExt, TryFutureExt};
use std::{
@@ -11,8 +10,6 @@ use std::{
},
};
use tokio::sync::Mutex;
use base64::Engine as _;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use guarden::guard;
@@ -20,7 +17,7 @@ use hmac::Mac;
use prost::Message;
use tokio::{
sync::broadcast,
sync::{Mutex, broadcast},
task::JoinSet,
time::{Duration, timeout},
};
@@ -36,6 +33,7 @@ use super::{
peer_session::{PeerSession, PeerSessionAction},
traffic_metrics::AggregateTrafficMetrics,
};
use crate::utils::BoxExt;
use crate::{
common::{
PeerId,
@@ -101,7 +99,7 @@ struct PeerSessionTunnelFilter {
enabled: bool,
my_peer_id: Arc<AtomicCell<PeerId>>,
peer_id: Arc<AtomicCell<Option<PeerId>>>,
session: Arc<ArcSwapOption<PeerSession>>,
session: Arc<std::sync::Mutex<Option<Arc<PeerSession>>>>,
}
impl PeerSessionTunnelFilter {
@@ -110,7 +108,7 @@ impl PeerSessionTunnelFilter {
enabled,
my_peer_id: Arc::new(AtomicCell::new(PeerId::default())),
peer_id: Arc::new(AtomicCell::new(None)),
session: Arc::new(ArcSwapOption::empty()),
session: Arc::new(std::sync::Mutex::new(None)),
}
}
@@ -119,7 +117,7 @@ impl PeerSessionTunnelFilter {
enabled,
my_peer_id: Arc::new(AtomicCell::new(my_peer_id)),
peer_id: Arc::new(AtomicCell::new(None)),
session: Arc::new(ArcSwapOption::empty()),
session: Arc::new(std::sync::Mutex::new(None)),
}
}
@@ -132,7 +130,7 @@ impl PeerSessionTunnelFilter {
}
fn set_session(&self, session: Arc<PeerSession>) {
self.session.store(Some(session));
*self.session.lock().unwrap() = Some(session);
}
fn should_skip_encrypt(&self, hdr: &crate::tunnel::packet_def::PeerManagerHeader) -> bool {
@@ -166,15 +164,16 @@ impl TunnelFilter for PeerSessionTunnelFilter {
return Some(data);
};
let mut guard = self.session.lock().unwrap();
let Some(session) = guard.as_mut() else {
return Some(data);
};
let my_peer_id = self.my_peer_id.load();
if my_peer_id != hdr.from_peer_id.get() || hdr.to_peer_id.get() != peer_id {
if my_peer_id != hdr.from_peer_id.get() {
return Some(data);
}
let session_guard = self.session.load();
let Some(session) = session_guard.as_deref() else {
return Some(data);
};
if let Err(e) = session.encrypt_payload(my_peer_id, peer_id, &mut data) {
tracing::warn!(
?my_peer_id,
@@ -219,8 +218,8 @@ impl TunnelFilter for PeerSessionTunnelFilter {
return Some(Ok(data));
}
let session_guard = self.session.load();
let Some(session) = session_guard.as_deref() else {
let mut guard = self.session.lock().unwrap();
let Some(session) = guard.as_mut() else {
return Some(Ok(data));
};
@@ -382,8 +381,7 @@ impl PeerConn {
noise_handshake_result: None,
tunnel: Arc::new(Mutex::new(
Box::new(guard!([mut mpsc_tunnel] mpsc_tunnel.close()))
as Box<dyn Any + Send + 'static>,
guard!([mut mpsc_tunnel] mpsc_tunnel.close()).boxed(),
)),
sink,
recv: Mutex::new(Some(recv)),
@@ -1645,45 +1643,6 @@ pub mod tests {
.unwrap_or(0)
}
#[test]
fn peer_session_filter_skips_relay_packet_for_next_hop() {
let my_peer_id = 10;
let next_hop_peer_id = 20;
let dst_peer_id = 30;
let filter = PeerSessionTunnelFilter::new_with_peer(my_peer_id, true);
filter.set_peer_id(next_hop_peer_id);
let session = Arc::new(PeerSession::new(
next_hop_peer_id,
PeerSession::new_root_key(),
1,
0,
"aes-gcm".to_string(),
"aes-gcm".to_string(),
None,
));
session.invalidate();
filter.set_session(session);
let mut packet = ZCPacket::new_with_payload(b"relay payload");
packet.fill_peer_manager_hdr(my_peer_id, dst_peer_id, PacketType::Data as u8);
packet
.mut_peer_manager_header()
.unwrap()
.set_encrypted(true);
let original_len = packet.buf_len();
let packet = filter
.before_send(packet)
.expect("relay packet should bypass next-hop session");
let hdr = packet.peer_manager_header().unwrap();
assert_eq!(hdr.from_peer_id.get(), my_peer_id);
assert_eq!(hdr.to_peer_id.get(), dst_peer_id);
assert!(hdr.is_encrypted());
assert_eq!(packet.buf_len(), original_len);
}
#[tokio::test]
async fn peer_conn_handshake_same_id() {
let ps = Arc::new(PeerSessionStore::new());
+1 -2
View File
@@ -6,7 +6,6 @@ use std::{
time::Duration,
};
use quanta::Instant;
use rand::{Rng, thread_rng};
use tokio::{
sync::broadcast,
@@ -178,7 +177,7 @@ impl PeerConnPinger {
sink.send(req).await?;
control_metrics.record_tx(req_len);
let now = Instant::now();
let now = std::time::Instant::now();
// wait until we get a pong packet in ctrl_resp_receiver
let resp = timeout(Duration::from_secs(2), async {
loop {
+12 -136
View File
@@ -2,18 +2,19 @@ use anyhow::Context;
use async_trait::async_trait;
use cidr::{Ipv4Cidr, Ipv6Cidr};
use dashmap::DashMap;
use quanta::Instant;
use std::collections::BTreeSet;
use std::{
fmt::Debug,
net::{IpAddr, Ipv4Addr, Ipv6Addr},
sync::{Arc, Weak, atomic::AtomicBool},
time::{Duration, SystemTime},
time::{Duration, Instant, SystemTime},
};
use tokio::sync::{Mutex, RwLock};
use tokio::{
sync::mpsc::{self, UnboundedReceiver, UnboundedSender},
sync::{
Mutex, RwLock,
mpsc::{self, UnboundedReceiver, UnboundedSender},
},
task::JoinSet,
};
@@ -1532,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
@@ -2196,13 +2184,14 @@ impl PeerManager {
#[cfg(test)]
mod tests {
use base64::Engine;
use std::{collections::HashMap, fmt::Debug, sync::Arc, time::Duration};
use quanta::Instant;
use std::{
fmt::Debug,
sync::Arc,
time::{Duration, Instant},
};
use crate::{
common::{
PeerId,
config::Flags,
global_ctx::{NetworkIdentity, tests::get_mock_global_ctx},
stats_manager::{LabelSet, LabelType, MetricName},
@@ -2217,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,
@@ -2261,16 +2250,6 @@ mod tests {
))
}
struct TestCostCalculator {
costs: HashMap<(PeerId, PeerId), i32>,
}
impl RouteCostCalculatorInterface for TestCostCalculator {
fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 {
*self.costs.get(&(src, dst)).unwrap_or(&1)
}
}
#[test]
fn recent_traffic_fanout_policy_only_marks_single_peer() {
assert!(PeerManager::should_mark_recent_traffic_for_fanout(0));
@@ -2678,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;
+8 -9
View File
@@ -6,7 +6,7 @@ use std::{
Arc, Weak,
atomic::{AtomicBool, AtomicU32, Ordering},
},
time::{Duration, SystemTime},
time::{Duration, Instant, SystemTime},
};
use arc_swap::ArcSwap;
@@ -24,7 +24,6 @@ use petgraph::{
use prefix_trie::PrefixMap;
use prost::Message;
use prost_reflect::{DynamicMessage, ReflectMessage};
use quanta::Instant;
use tokio::{
select,
sync::Mutex,
@@ -2169,9 +2168,9 @@ struct PeerRouteServiceImpl {
interface_peers_generation: AtomicU64,
applied_interface_peers_generation: AtomicU64,
last_update_my_foreign_network: AtomicCell<Option<Instant>>,
last_update_my_foreign_network: AtomicCell<Option<std::time::Instant>>,
peer_info_last_update: AtomicCell<Instant>,
peer_info_last_update: AtomicCell<std::time::Instant>,
}
impl Debug for PeerRouteServiceImpl {
@@ -2238,7 +2237,7 @@ impl PeerRouteServiceImpl {
last_update_my_foreign_network: AtomicCell::new(None),
peer_info_last_update: AtomicCell::new(Instant::now()),
peer_info_last_update: AtomicCell::new(std::time::Instant::now()),
}
}
@@ -2434,7 +2433,7 @@ impl PeerRouteServiceImpl {
}
self.last_update_my_foreign_network
.store(Some(Instant::now()));
.store(Some(std::time::Instant::now()));
let foreign_networks = self
.interface
@@ -3155,12 +3154,12 @@ impl PeerRouteServiceImpl {
"update_peer_info_last_update, my_peer_id: {:?}, prev: {:?}, new: {:?}",
self.my_peer_id,
self.peer_info_last_update.load(),
Instant::now()
std::time::Instant::now()
);
self.peer_info_last_update.store(Instant::now());
self.peer_info_last_update.store(std::time::Instant::now());
}
fn get_peer_info_last_update(&self) -> Instant {
fn get_peer_info_last_update(&self) -> std::time::Instant {
self.peer_info_last_update.load()
}
+15 -167
View File
@@ -1,7 +1,7 @@
use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::net::SocketAddr;
use crate::{
common::{global_ctx::ArcGlobalCtx, network::IPCollector},
common::global_ctx::ArcGlobalCtx,
proto::{
common::Void,
peer_rpc::{
@@ -12,8 +12,6 @@ use crate::{
tunnel::udp,
};
const MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16;
fn remove_easytier_managed_ipv6s(ret: &mut GetIpListResponse, global_ctx: &ArcGlobalCtx) {
ret.interface_ipv6s.retain(|ip| {
let ip = std::net::Ipv6Addr::from(*ip);
@@ -30,86 +28,6 @@ fn remove_easytier_managed_ipv6s(ret: &mut GetIpListResponse, global_ctx: &ArcGl
}
}
fn is_usable_preferred_src_ipv6(ip: &Ipv6Addr, global_ctx: &ArcGlobalCtx) -> bool {
!global_ctx.is_ip_easytier_managed_ipv6(ip)
&& !ip.is_loopback()
&& !ip.is_unspecified()
&& !ip.is_unique_local()
&& !ip.is_unicast_link_local()
&& !ip.is_multicast()
}
async fn local_preferred_src_ipv6(
global_ctx: &ArcGlobalCtx,
preferred_src_ipv6: Option<crate::proto::common::Ipv6Addr>,
) -> Option<udp::PreferredIpv6Source> {
let preferred_src_ipv6 = preferred_src_ipv6.map(Ipv6Addr::from)?;
if !is_usable_preferred_src_ipv6(&preferred_src_ipv6, global_ctx) {
tracing::debug!(
?preferred_src_ipv6,
"ignore unusable preferred IPv6 source for udp hole punch"
);
return None;
}
let ifaces = IPCollector::collect_interfaces(global_ctx.net_ns.clone(), false).await;
for iface in ifaces {
let is_local = iface.ips.iter().any(|ip| match ip.ip() {
IpAddr::V6(v6) => v6 == preferred_src_ipv6,
IpAddr::V4(_) => false,
});
if is_local {
tracing::debug!(
?preferred_src_ipv6,
ifindex = iface.index,
"use preferred IPv6 source for udp hole punch"
);
return Some(udp::PreferredIpv6Source {
ip: preferred_src_ipv6,
ifindex: iface.index,
});
}
}
tracing::debug!(
?preferred_src_ipv6,
"ignore non-local preferred IPv6 source for udp hole punch"
);
None
}
fn connector_addrs_from_request(
req: SendUdpHolePunchPacketRequest,
) -> rpc_types::error::Result<(u16, Vec<SocketAddr>, Option<crate::proto::common::Ipv6Addr>)> {
let listener_port = u16::try_from(req.listener_port)
.map_err(|_| anyhow::anyhow!("listener_port is out of range: {}", req.listener_port))?;
let mut connector_addrs = req
.connector_addrs
.into_iter()
.map(SocketAddr::from)
.collect::<Vec<_>>();
if connector_addrs.is_empty() {
connector_addrs.push(
req.connector_addr
.ok_or(anyhow::anyhow!("connector_addr is required"))?
.into(),
);
}
let mut deduped = Vec::with_capacity(connector_addrs.len());
for addr in connector_addrs {
if !deduped.contains(&addr) {
deduped.push(addr);
}
if deduped.len() >= MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS {
break;
}
}
Ok((listener_port, deduped, req.preferred_src_ipv6))
}
#[derive(Clone)]
pub struct DirectConnectorManagerRpcServer {
// TODO: this only cache for one src peer, should make it global
@@ -149,38 +67,23 @@ impl DirectConnectorRpc for DirectConnectorManagerRpcServer {
_: BaseController,
req: SendUdpHolePunchPacketRequest,
) -> rpc_types::error::Result<Void> {
let (listener_port, connector_addrs, preferred_src_ipv6) =
connector_addrs_from_request(req)?;
let preferred_src_ipv6 =
local_preferred_src_ipv6(&self.global_ctx, preferred_src_ipv6).await;
let listener_port = req.listener_port as u16;
let connector_addr: SocketAddr = req
.connector_addr
.ok_or(anyhow::anyhow!("connector_addr is required"))?
.into();
tracing::info!(
?connector_addrs,
?preferred_src_ipv6,
listener_port,
"Sending udp hole punch packet"
"Sending udp hole punch packet to {} from listener port {}",
connector_addr,
listener_port
);
// send 3 packets to the connector
for _ in 0..3 {
for connector_addr in &connector_addrs {
let ret = match connector_addr {
SocketAddr::V4(addr) => {
udp::send_v4_hole_punch_packet(listener_port, *addr).await
}
SocketAddr::V6(addr) => {
udp::send_v6_hole_punch_packet(listener_port, *addr, preferred_src_ipv6)
.await
}
};
if let Err(e) = ret {
tracing::debug!(
?e,
?connector_addr,
listener_port,
"send udp hole punch packet failed"
);
}
match connector_addr {
SocketAddr::V4(addr) => udp::send_v4_hole_punch_packet(listener_port, addr).await?,
SocketAddr::V6(addr) => udp::send_v6_hole_punch_packet(listener_port, addr).await?,
}
tokio::time::sleep(std::time::Duration::from_millis(30)).await;
}
@@ -196,12 +99,11 @@ impl DirectConnectorManagerRpcServer {
#[cfg(test)]
mod tests {
use std::{collections::BTreeSet, net::SocketAddr};
use std::collections::BTreeSet;
use crate::{
common::global_ctx::tests::get_mock_global_ctx,
peers::peer_rpc_service::{connector_addrs_from_request, remove_easytier_managed_ipv6s},
proto::peer_rpc::{GetIpListResponse, SendUdpHolePunchPacketRequest},
peers::peer_rpc_service::remove_easytier_managed_ipv6s, proto::peer_rpc::GetIpListResponse,
};
#[tokio::test]
@@ -231,58 +133,4 @@ mod tests {
assert_eq!(ip_list.public_ipv6, None);
assert_eq!(ip_list.interface_ipv6s, vec![physical_ipv6.into()]);
}
#[test]
fn hole_punch_request_prefers_batch_connector_addrs() {
let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap();
let first_batch_addr: SocketAddr = "[2001:db8::2]:10002".parse().unwrap();
let second_batch_addr: SocketAddr = "[2001:db8::3]:10003".parse().unwrap();
let preferred_src_ipv6: std::net::Ipv6Addr = "2001:db8::4".parse().unwrap();
let (listener_port, connector_addrs, preferred_src) =
connector_addrs_from_request(SendUdpHolePunchPacketRequest {
connector_addr: Some(old_addr.into()),
listener_port: 11010,
preferred_src_ipv6: Some(preferred_src_ipv6.into()),
connector_addrs: vec![
first_batch_addr.into(),
first_batch_addr.into(),
second_batch_addr.into(),
],
})
.unwrap();
assert_eq!(listener_port, 11010);
assert_eq!(connector_addrs, vec![first_batch_addr, second_batch_addr]);
assert_eq!(preferred_src, Some(preferred_src_ipv6.into()));
}
#[test]
fn hole_punch_request_falls_back_to_legacy_connector_addr() {
let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap();
let (_, connector_addrs, _) = connector_addrs_from_request(SendUdpHolePunchPacketRequest {
connector_addr: Some(old_addr.into()),
listener_port: 11010,
preferred_src_ipv6: None,
connector_addrs: vec![],
})
.unwrap();
assert_eq!(connector_addrs, vec![old_addr]);
}
#[test]
fn hole_punch_request_rejects_out_of_range_listener_port() {
let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap();
let ret = connector_addrs_from_request(SendUdpHolePunchPacketRequest {
connector_addr: Some(old_addr.into()),
listener_port: u16::MAX as u32 + 1,
preferred_src_ipv6: None,
connector_addrs: vec![],
});
assert!(ret.is_err());
}
}
+17 -129
View File
@@ -2,12 +2,9 @@ use std::sync::{
Arc, RwLock,
atomic::{AtomicBool, Ordering},
};
use std::time::Duration;
use anyhow::anyhow;
use crossbeam::atomic::AtomicCell;
use dashmap::DashMap;
use quanta::Instant;
use super::secure_datagram::{SecureDatagramDirection, SecureDatagramSession};
use crate::{
@@ -15,8 +12,6 @@ use crate::{
tunnel::packet_def::ZCPacket,
};
const SESSION_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
pub struct UpsertResponderSessionReturn {
pub session: Arc<PeerSession>,
pub action: PeerSessionAction,
@@ -49,25 +44,7 @@ impl SessionKey {
#[derive(Clone)]
pub struct PeerSessionStore {
sessions: Arc<DashMap<SessionKey, PeerSessionEntry>>,
}
struct PeerSessionEntry {
session: Arc<PeerSession>,
last_used_at: AtomicCell<Instant>,
}
impl PeerSessionEntry {
fn new(session: Arc<PeerSession>) -> Self {
Self {
session,
last_used_at: AtomicCell::new(Instant::now()),
}
}
fn touch(&self) {
self.last_used_at.store(Instant::now());
}
sessions: Arc<DashMap<SessionKey, Arc<PeerSession>>>,
}
impl Default for PeerSessionStore {
@@ -84,11 +61,7 @@ impl PeerSessionStore {
}
pub fn get(&self, key: &SessionKey) -> Option<Arc<PeerSession>> {
let session = {
let entry = self.sessions.get(key)?;
entry.touch();
entry.session.clone()
};
let session = self.sessions.get(key)?.clone();
if session.is_valid() {
Some(session)
} else {
@@ -102,20 +75,12 @@ impl PeerSessionStore {
}
pub fn insert_session(&self, key: SessionKey, session: Arc<PeerSession>) {
self.sessions.insert(key, PeerSessionEntry::new(session));
self.sessions.insert(key, session);
}
pub fn evict_unused_sessions(&self) {
self.evict_unused_sessions_idle(SESSION_IDLE_TIMEOUT);
}
pub fn evict_unused_sessions_idle(&self, idle: Duration) {
let now = Instant::now();
self.sessions.retain(|_key, entry| {
entry.session.is_valid()
&& (Arc::strong_count(&entry.session) > 1
|| now.saturating_duration_since(entry.last_used_at.load()) < idle)
});
self.sessions
.retain(|_key, session| Arc::strong_count(session) > 1);
shrink_dashmap(&self.sessions, None);
}
@@ -128,14 +93,11 @@ impl PeerSessionStore {
recv_algorithm: String,
peer_static_pubkey: Option<[u8; 32]>,
) -> Result<UpsertResponderSessionReturn, anyhow::Error> {
tracing::event!(tracing::Level::INFO, ?key, "upsert_responder_session");
tracing::event!(tracing::Level::INFO, "upsert_responder_session {:?}", key);
let existing = self
.sessions
.get(key)
.map(|v| {
v.touch();
v.session.clone()
})
.map(|v| v.clone())
.filter(|s| s.is_valid());
match existing {
None => {
@@ -151,8 +113,7 @@ impl PeerSessionStore {
recv_algorithm,
peer_static_pubkey,
));
self.sessions
.insert(key.clone(), PeerSessionEntry::new(session.clone()));
self.sessions.insert(key.clone(), session.clone());
Ok(UpsertResponderSessionReturn {
session,
action: PeerSessionAction::Create,
@@ -217,14 +178,16 @@ impl PeerSessionStore {
PeerSessionAction::Sync | PeerSessionAction::Create => {
let root_key = root_key_32.ok_or_else(|| anyhow!("missing root_key"))?;
if let Some(existing) = self.sessions.get(key)
&& !existing.session.is_valid()
&& !existing.is_valid()
{
drop(existing);
self.sessions.remove(key);
}
let session = {
let entry = self.sessions.entry(key.clone()).or_insert_with(|| {
PeerSessionEntry::new(Arc::new(PeerSession::new(
let session = self
.sessions
.entry(key.clone())
.or_insert_with(|| {
Arc::new(PeerSession::new(
key.peer_id,
root_key,
b_session_generation,
@@ -232,11 +195,9 @@ impl PeerSessionStore {
send_algorithm.clone(),
recv_algorithm.clone(),
peer_static_pubkey,
)))
});
entry.touch();
entry.session.clone()
};
))
})
.clone();
session.check_encrypt_algo_same(&send_algorithm, &recv_algorithm)?;
session.check_or_set_peer_static_pubkey(peer_static_pubkey)?;
session.sync_root_key(
@@ -458,77 +419,4 @@ mod tests {
SecureDatagramSession::SYNC_RX_GRACE_AFTER_MS
);
}
#[test]
fn peer_session_store_keeps_recent_session_without_external_refs() {
let store = PeerSessionStore::new();
let key = SessionKey::new("net".to_string(), 20);
let session = Arc::new(PeerSession::new(
20,
PeerSession::new_root_key(),
1,
0,
"aes-gcm".to_string(),
"aes-gcm".to_string(),
None,
));
store.insert_session(key.clone(), session);
assert!(store.get(&key).is_some());
store.evict_unused_sessions();
assert!(
store.get(&key).is_some(),
"recent relay sessions should survive the periodic GC"
);
}
#[test]
fn peer_session_store_evicts_idle_session_without_external_refs() {
let store = PeerSessionStore::new();
let key = SessionKey::new("net".to_string(), 20);
let session = Arc::new(PeerSession::new(
20,
PeerSession::new_root_key(),
1,
0,
"aes-gcm".to_string(),
"aes-gcm".to_string(),
None,
));
store.insert_session(key.clone(), session);
store.evict_unused_sessions_idle(Duration::from_millis(0));
assert!(
store.get(&key).is_none(),
"idle sessions without external users should still be collected"
);
}
#[test]
fn peer_session_store_evicts_invalid_recent_session() {
let store = PeerSessionStore::new();
let key = SessionKey::new("net".to_string(), 20);
let session = Arc::new(PeerSession::new(
20,
PeerSession::new_root_key(),
1,
0,
"aes-gcm".to_string(),
"aes-gcm".to_string(),
None,
));
store.insert_session(key.clone(), session);
let session = store.get(&key).unwrap();
session.invalidate();
drop(session);
store.evict_unused_sessions();
assert!(
!store.sessions.contains_key(&key),
"invalid sessions should not be kept by recent activity"
);
}
}

Some files were not shown because too many files have changed in this diff Show More