From cc2d58bb00fd58be1ac78ad82cd18dfb410fdabd Mon Sep 17 00:00:00 2001 From: KKRainbow <5665404+KKRainbow@users.noreply.github.com> Date: Wed, 23 Sep 2026 21:43:05 +0800 Subject: [PATCH] Add shared virtual NIC core with per-member routing Share a named TUN across instances while keeping unnamed devices dedicated. Track member addresses, routes, and MTU, and reject overlapping destinations except the common Magic DNS host route. Dispatch packets between the physical TUN and member tunnels. Keep IPv4 and ordinary IPv6 source translation for mobile wrong-source traffic. Update platform configuration and runtime host integration, with unit and network namespace coverage for routing and lifecycle. Add the packet dependency and keep core CI matrix jobs independent. --- .github/workflows/core.yml | 2 +- Cargo.lock | 34 + easytier/Cargo.toml | 3 +- easytier/src/common/ifcfg/darwin.rs | 76 +- easytier/src/common/ifcfg/mod.rs | 21 + easytier/src/common/ifcfg/netlink.rs | 115 +- easytier/src/common/ifcfg/netlink_wire.rs | 24 + easytier/src/common/ifcfg/windows.rs | 1 + easytier/src/instance/composition.rs | 9 +- easytier/src/instance/dns_server/runner.rs | 121 +- .../instance/dns_server/server_instance.rs | 83 +- easytier/src/instance/factory.rs | 66 + easytier/src/instance/mod.rs | 3 + easytier/src/instance/runtime_host.rs | 49 +- .../src/instance/runtime_host/magic_dns.rs | 7 +- .../src/instance/runtime_host/tun_common.rs | 36 +- .../src/instance/runtime_host/tun_desktop.rs | 49 +- .../src/instance/runtime_host/tun_mobile.rs | 116 +- easytier/src/instance/shared_virtual_nic.rs | 1738 +++++++++++ .../instance/shared_virtual_nic/dispatcher.rs | 2657 +++++++++++++++++ easytier/src/instance/test_instance.rs | 45 +- easytier/src/instance/virtual_nic.rs | 759 +++-- easytier/src/tests/mod.rs | 3 + easytier/src/tests/shared_virtual_nic.rs | 476 +++ easytier/src/tests/three_node.rs | 151 + 25 files changed, 6357 insertions(+), 287 deletions(-) create mode 100644 easytier/src/instance/shared_virtual_nic.rs create mode 100644 easytier/src/instance/shared_virtual_nic/dispatcher.rs create mode 100644 easytier/src/tests/shared_virtual_nic.rs diff --git a/.github/workflows/core.yml b/.github/workflows/core.yml index 93634d72..0bbb4479 100644 --- a/.github/workflows/core.yml +++ b/.github/workflows/core.yml @@ -82,7 +82,7 @@ jobs: easytier-web/frontend/dist/* build: strategy: - fail-fast: true + fail-fast: false matrix: include: - TARGET: x86_64-unknown-linux-musl diff --git a/Cargo.lock b/Cargo.lock index db9dfec8..4b4e447b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2594,6 +2594,7 @@ dependencies = [ "percent-encoding", "pin-project-lite", "pnet_datalink", + "pnet_packet", "prost 0.14.4", "quanta", "quinn", @@ -6911,6 +6912,39 @@ dependencies = [ "winapi", ] +[[package]] +name = "pnet_macros" +version = "0.35.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13325ac86ee1a80a480b0bc8e3d30c25d133616112bb16e86f712dcf8a71c863" +dependencies = [ + "proc-macro2", + "quote", + "regex", + "syn 2.0.119", +] + +[[package]] +name = "pnet_macros_support" +version = "0.35.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eed67a952585d509dd0003049b1fc56b982ac665c8299b124b90ea2bdb3134ab" +dependencies = [ + "pnet_base", +] + +[[package]] +name = "pnet_packet" +version = "0.35.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c96ebadfab635fcc23036ba30a7d33a80c39e8461b8bd7dc7bb186acb96560f" +dependencies = [ + "glob", + "pnet_base", + "pnet_macros", + "pnet_macros_support", +] + [[package]] name = "pnet_sys" version = "0.35.0" diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index d04a3a3d..e319968b 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -153,6 +153,7 @@ rand.workspace = true serde = { workspace = true, features = ["derive"] } pnet_datalink = { version = "0.35.0", optional = true } +pnet_packet = "0.35.0" smoltcp = { workspace = true, optional = true, features = [ "std", "medium-ethernet", @@ -362,7 +363,7 @@ mimalloc = ["dep:mimalloc"] aes-gcm = ["easytier-core/aes-gcm"] openssl-crypto = ["easytier-core/openssl-crypto"] ring-crypto = ["easytier-core/ring-crypto"] -tun = ["dep:tun", "linux-netlink"] +tun = ["dep:tun", "linux-netlink", "tokio/rt-multi-thread"] linux-netlink = ["dep:netlink-sys"] proxy-cidr-monitor = ["easytier-core/proxy-cidr-monitor"] websocket = [ diff --git a/easytier/src/common/ifcfg/darwin.rs b/easytier/src/common/ifcfg/darwin.rs index 0412d9e2..41f939f9 100644 --- a/easytier/src/common/ifcfg/darwin.rs +++ b/easytier/src/common/ifcfg/darwin.rs @@ -1,10 +1,38 @@ -use std::net::Ipv4Addr; +use std::{collections::BTreeSet, net::Ipv4Addr}; use super::{Error, IfConfiguerTrait, cidr_to_subnet_mask, run_shell_cmd}; use async_trait::async_trait; use cidr::{Ipv4Inet, Ipv6Inet}; +use tokio::sync::Mutex; + +#[derive(Default)] +pub struct MacIfConfiger { + configured_ipv4: Mutex>, +} + +impl MacIfConfiger { + fn build_add_ipv4_cmd(name: &str, addr: Ipv4Inet, has_configured_ipv4: bool) -> String { + let address = addr.address(); + if has_configured_ipv4 { + format!( + "ifconfig {} alias {:?} {:?} netmask {}", + name, + address, + address, + cidr_to_subnet_mask(addr.network_length()) + ) + } else { + format!( + "ifconfig {} {:?}/{:?} {:?} up", + name, + address, + addr.network_length(), + address, + ) + } + } +} -pub struct MacIfConfiger {} #[async_trait] impl IfConfiguerTrait for MacIfConfiger { async fn add_ipv4_route( @@ -14,12 +42,28 @@ impl IfConfiguerTrait for MacIfConfiger { cidr_prefix: u8, cost: Option, ) -> Result<(), Error> { + self.add_ipv4_route_with_source_hint(name, address, cidr_prefix, cost, None) + .await + } + + async fn add_ipv4_route_with_source_hint( + &self, + name: &str, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + source_hint: Option, + ) -> Result<(), Error> { + let source_hint = source_hint + .map(|source| format!(" -ifa {}", source)) + .unwrap_or_default(); run_shell_cmd( format!( - "route -n add {} -netmask {} -interface {} -hopcount {}", + "route -n add {} -netmask {} -interface {}{} -hopcount {}", address, cidr_to_subnet_mask(cidr_prefix), name, + source_hint, cost.unwrap_or(7) ) .as_str(), @@ -51,14 +95,15 @@ impl IfConfiguerTrait for MacIfConfiger { address: Ipv4Addr, cidr_prefix: u8, ) -> Result<(), Error> { - run_shell_cmd( - format!( - "ifconfig {} {:?}/{:?} {:?} up", - name, address, cidr_prefix, address, - ) - .as_str(), - ) - .await + let addr = Ipv4Inet::new(address, cidr_prefix).map_err(|err| { + anyhow::anyhow!("invalid IPv4 address {address}/{cidr_prefix}: {err:?}") + })?; + let mut configured_ipv4 = self.configured_ipv4.lock().await; + let cmd = Self::build_add_ipv4_cmd(name, addr, !configured_ipv4.is_empty()); + + run_shell_cmd(cmd.as_str()).await?; + configured_ipv4.insert(addr); + Ok(()) } async fn set_link_status(&self, name: &str, up: bool) -> Result<(), Error> { @@ -67,11 +112,16 @@ impl IfConfiguerTrait for MacIfConfiger { } async fn remove_ip(&self, name: &str, ip: Option) -> Result<(), Error> { + let mut configured_ipv4 = self.configured_ipv4.lock().await; if let Some(ip) = ip { - run_shell_cmd(format!("ifconfig {} inet {} delete", name, ip.address()).as_str()).await + run_shell_cmd(format!("ifconfig {} inet {} delete", name, ip.address()).as_str()) + .await?; + configured_ipv4.remove(&ip); } else { - run_shell_cmd(format!("ifconfig {} inet delete", name).as_str()).await + run_shell_cmd(format!("ifconfig {} inet delete", name).as_str()).await?; + configured_ipv4.clear(); } + Ok(()) } async fn set_mtu(&self, name: &str, mtu: u32) -> Result<(), Error> { diff --git a/easytier/src/common/ifcfg/mod.rs b/easytier/src/common/ifcfg/mod.rs index 0a64a33c..da3a55e1 100644 --- a/easytier/src/common/ifcfg/mod.rs +++ b/easytier/src/common/ifcfg/mod.rs @@ -40,6 +40,16 @@ pub trait IfConfiguerTrait: Send + Sync { ) -> Result<(), Error> { Ok(()) } + async fn add_ipv4_route_with_source_hint( + &self, + name: &str, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + _source_hint: Option, + ) -> Result<(), Error> { + self.add_ipv4_route(name, address, cidr_prefix, cost).await + } async fn remove_ipv4_route( &self, _name: &str, @@ -48,6 +58,16 @@ pub trait IfConfiguerTrait: Send + Sync { ) -> Result<(), Error> { Ok(()) } + async fn remove_ipv4_route_with_cost_and_source_hint( + &self, + name: &str, + address: Ipv4Addr, + cidr_prefix: u8, + _cost: Option, + _source_hint: Option, + ) -> Result<(), Error> { + self.remove_ipv4_route(name, address, cidr_prefix).await + } async fn add_ipv4_ip( &self, _name: &str, @@ -157,6 +177,7 @@ async fn run_shell_cmd(cmd: &str) -> Result<(), Error> { Ok(()) } +#[derive(Default)] pub struct DummyIfConfiger {} #[async_trait] impl IfConfiguerTrait for DummyIfConfiger {} diff --git a/easytier/src/common/ifcfg/netlink.rs b/easytier/src/common/ifcfg/netlink.rs index a401e744..d5068325 100644 --- a/easytier/src/common/ifcfg/netlink.rs +++ b/easytier/src/common/ifcfg/netlink.rs @@ -133,6 +133,7 @@ fn dump_netlink_messages( receive_netlink_dump(builder) } +#[derive(Default)] pub struct NetlinkIfConfiger {} impl NetlinkIfConfiger { @@ -335,6 +336,27 @@ impl NetlinkIfConfiger { }) .collect()) } + + fn ipv4_route_message( + ifindex: u32, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + source_hint: Option, + ) -> RouteMessage { + let mut builder = RouteMessageBuilder::new(libc::AF_INET as u8) + .destination(IpAddr::V4(address), cidr_prefix) + .oif(ifindex) + .priority(cost.unwrap_or(65535) as u32) + .table(libc::RT_TABLE_MAIN.into()) + .static_protocol() + .universe_scope() + .route_type(RouteType::Unicast); + if let Some(source_hint) = source_hint { + builder = builder.preferred_source(IpAddr::V4(source_hint)); + } + builder.build() + } } #[async_trait] @@ -346,15 +368,25 @@ impl IfConfiguerTrait for NetlinkIfConfiger { cidr_prefix: u8, cost: Option, ) -> Result<(), Error> { - let message = RouteMessageBuilder::new(libc::AF_INET as u8) - .destination(IpAddr::V4(address), cidr_prefix) - .oif(Self::get_interface_index(name)?) - .priority(cost.unwrap_or(65535) as u32) - .table(libc::RT_TABLE_MAIN.into()) - .static_protocol() - .universe_scope() - .route_type(RouteType::Unicast) - .build(); + self.add_ipv4_route_with_source_hint(name, address, cidr_prefix, cost, None) + .await + } + + async fn add_ipv4_route_with_source_hint( + &self, + name: &str, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + source_hint: Option, + ) -> Result<(), Error> { + let message = NetlinkIfConfiger::ipv4_route_message( + NetlinkIfConfiger::get_interface_index(name)?, + address, + cidr_prefix, + cost, + source_hint, + ); let request = message_request( RTM_NEWROUTE, NLM_F_ACK | NLM_F_CREATE | NLM_F_EXCL | NLM_F_REQUEST, @@ -390,6 +422,28 @@ impl IfConfiguerTrait for NetlinkIfConfiger { Ok(()) } + async fn remove_ipv4_route_with_cost_and_source_hint( + &self, + name: &str, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + source_hint: Option, + ) -> Result<(), Error> { + let message = Self::ipv4_route_message( + Self::get_interface_index(name)?, + address, + cidr_prefix, + cost, + source_hint, + ); + let request = message_request(RTM_DELROUTE, NLM_F_ACK | NLM_F_REQUEST, &message)?; + match send_netlink_req_and_wait_ack(request) { + Err(Error::IOError(err)) if err.raw_os_error() == Some(libc::ESRCH) => Ok(()), + result => result, + } + } + async fn add_ipv4_ip( &self, name: &str, @@ -696,6 +750,49 @@ mod tests { assert!(!routes.contains(&IpAddr::V4("10.5.5.0".parse().unwrap()))); } + #[serial_test::serial] + #[tokio::test] + async fn remove_ipv4_route_with_source_hint_keeps_other_metric() { + let iface = test_iface_name("rm"); + let _link = ScopedDummyLink::new(&iface); + let ifcfg = NetlinkIfConfiger {}; + let address = "10.231.1.1".parse().unwrap(); + let destination = "10.99.0.0".parse().unwrap(); + + ifcfg.add_ipv4_ip(&iface, address, 24).await.unwrap(); + for cost in [123, 124] { + ifcfg + .add_ipv4_route_with_source_hint(&iface, destination, 24, Some(cost), Some(address)) + .await + .unwrap(); + } + + ifcfg + .remove_ipv4_route_with_cost_and_source_hint( + &iface, + destination, + 24, + Some(123), + Some(address), + ) + .await + .unwrap(); + let routes = run_cmd(&format!("ip -4 route show 10.99.0.0/24 dev {iface}")); + assert!(!routes.contains("metric 123")); + assert!(routes.contains("metric 124")); + + ifcfg + .remove_ipv4_route_with_cost_and_source_hint( + &iface, + destination, + 24, + Some(123), + Some(address), + ) + .await + .unwrap(); + } + #[serial_test::serial] #[tokio::test] async fn ipv6_addr_readback_test() { diff --git a/easytier/src/common/ifcfg/netlink_wire.rs b/easytier/src/common/ifcfg/netlink_wire.rs index 6bd92925..7dbb7284 100644 --- a/easytier/src/common/ifcfg/netlink_wire.rs +++ b/easytier/src/common/ifcfg/netlink_wire.rs @@ -37,6 +37,7 @@ const RTA_DST: u16 = 1; const RTA_SRC: u16 = 2; const RTA_OIF: u16 = 4; const RTA_PRIORITY: u16 = 6; +const RTA_PREFSRC: u16 = 7; const RTA_TABLE: u16 = 15; const NDA_DST: u16 = 1; @@ -488,6 +489,13 @@ impl RouteMessageBuilder { self } + pub(crate) fn preferred_source(mut self, address: IpAddr) -> Self { + self.message + .attributes + .push(Attribute::new(RTA_PREFSRC, ip_bytes(address))); + self + } + pub(crate) fn table(mut self, table: u32) -> Self { if let Ok(table) = u8::try_from(table) { self.message.table = table; @@ -632,6 +640,22 @@ mod tests { assert_eq!(encode(&decoded), bytes); } + #[test] + fn route_builder_encodes_preferred_source() { + let source = "10.231.1.1".parse().unwrap(); + let message = RouteMessageBuilder::new(libc::AF_INET as u8) + .destination("10.99.0.0".parse().unwrap(), 24) + .preferred_source(source) + .oif(7) + .table(libc::RT_TABLE_MAIN.into()) + .build(); + let bytes = encode(&message); + let attributes = parse_attributes(&bytes[12..]).unwrap(); + assert!(attributes.iter().any(|attribute| { + attribute.kind == RTA_PREFSRC && attribute.value == ip_bytes(source) + })); + } + #[test] fn route_parser_reads_ipv6_source_prefix() { let mut bytes = vec![ diff --git a/easytier/src/common/ifcfg/windows.rs b/easytier/src/common/ifcfg/windows.rs index 15e4b23d..e2ddb129 100644 --- a/easytier/src/common/ifcfg/windows.rs +++ b/easytier/src/common/ifcfg/windows.rs @@ -22,6 +22,7 @@ use winreg::{ }; use super::{Error, IfConfiguerTrait}; +#[derive(Default)] pub struct WindowsIfConfiger {} fn format_win_error(error: u32) -> String { diff --git a/easytier/src/instance/composition.rs b/easytier/src/instance/composition.rs index e8484ff8..0241c505 100644 --- a/easytier/src/instance/composition.rs +++ b/easytier/src/instance/composition.rs @@ -32,6 +32,8 @@ use crate::{ }; use super::host::{NativeInstanceHost, native_instance_host}; +#[cfg(feature = "tun")] +use super::shared_virtual_nic::ArcSharedVirtualNicRegistry; #[cfg(feature = "kcp")] use crate::gateway::kcp_proxy::KcpProxyService; #[cfg(feature = "quic")] @@ -45,6 +47,7 @@ pub(crate) type NativeCoreInstance = CoreInstance; pub(crate) fn compose_native_core_instance( toml_config: TomlConfig, process_runtime: Arc, + #[cfg(feature = "tun")] shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, compact_runtime: bool, ) -> anyhow::Result> { let host_config = if compact_runtime { @@ -58,7 +61,11 @@ pub(crate) fn compose_native_core_instance( &normalized, &host_config, )); - let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone()); + let runtime_host = NativeInstanceRuntimeHost::new( + global_ctx.clone(), + #[cfg(feature = "tun")] + shared_virtual_nic_registry, + ); let mut adapters = runtime_core_host_adapters_with_packet_egress_and_config( global_ctx.clone(), process_runtime, diff --git a/easytier/src/instance/dns_server/runner.rs b/easytier/src/instance/dns_server/runner.rs index 76791c68..f47b5094 100644 --- a/easytier/src/instance/dns_server/runner.rs +++ b/easytier/src/instance/dns_server/runner.rs @@ -5,7 +5,15 @@ use std::{net::Ipv4Addr, sync::Arc, time::Duration}; use easytier_core::instance::CorePacketPlane; -use crate::common::global_ctx::ArcGlobalCtx; +use crate::{ + common::{ + error::Error as EtError, + global_ctx::ArcGlobalCtx, + ifcfg::{IfConfiger, IfConfiguerTrait}, + netns::NetNS, + }, + instance::virtual_nic::NicBackend, +}; use super::{client_instance::MagicDnsClientInstance, server_instance::MagicDnsServerInstance}; @@ -17,6 +25,52 @@ pub struct DnsRunner { tun_dev: Option, tun_inet: Ipv4Inet, fake_ip: Ipv4Addr, + shared_route_backend: Option, +} + +#[derive(Clone)] +struct MagicDnsFakeIpRouteClaim { + tun_dev: Option, + net_ns: NetNS, + fake_ip: Ipv4Addr, + route_backend: NicBackend, +} + +impl MagicDnsFakeIpRouteClaim { + async fn add(&self) -> anyhow::Result<()> { + let cost = if cfg!(target_os = "windows") { + Some(4) + } else { + None + }; + + match self + .route_backend + .add_route_with_cost(self.fake_ip, 32, cost) + .await + { + Err(EtError::IOError(err)) + if err.kind() == std::io::ErrorKind::AlreadyExists && self.tun_dev.is_some() => + { + let ifcfg = IfConfiger::default(); + let _guard = self.net_ns.guard(); + ifcfg + .remove_ipv4_route(self.tun_dev.as_deref().unwrap(), self.fake_ip, 32) + .await?; + self.route_backend + .add_route_with_cost(self.fake_ip, 32, cost) + .await?; + Ok(()) + } + result => result.map_err(Into::into), + } + } + + async fn remove(&self) { + if let Err(err) = self.route_backend.remove_route(self.fake_ip, 32).await { + tracing::warn!(?err, fake_ip = ?self.fake_ip, "remove magic dns route failed"); + } + } } impl DnsRunner { @@ -35,9 +89,15 @@ impl DnsRunner { tun_dev, tun_inet, fake_ip, + shared_route_backend: None, } } + pub(crate) fn with_shared_route_backend(mut self, route_backend: Option) -> Self { + self.shared_route_backend = route_backend; + self + } + async fn clean_env(&mut self) { if let Some(server) = self.server.take() { server.clean_env().await; @@ -45,17 +105,53 @@ impl DnsRunner { self.client.take(); } + fn should_manage_fake_ip_route_externally(&self) -> bool { + self.shared_route_backend.is_some() && !self.tun_inet.contains(&self.fake_ip) + } + + fn fake_ip_route_claim(&self) -> Option { + if !self.should_manage_fake_ip_route_externally() { + return None; + } + + Some(MagicDnsFakeIpRouteClaim { + tun_dev: self.tun_dev.clone(), + net_ns: self.global_ctx.net_ns.clone(), + fake_ip: self.fake_ip, + route_backend: self.shared_route_backend.clone()?, + }) + } + async fn run_once(&mut self) -> anyhow::Result<()> { + if let Some(claim) = self.fake_ip_route_claim() { + claim + .add() + .await + .map_err(|err| anyhow::anyhow!("failed to add magic dns fake-ip route: {err}"))?; + } + // try server first - match MagicDnsServerInstance::new( - self.packet_plane.clone(), - self.global_ctx.clone(), - self.tun_dev.clone(), - self.tun_inet, - self.fake_ip, - ) - .await - { + let server_result = if self.should_manage_fake_ip_route_externally() { + MagicDnsServerInstance::new_with_external_fake_ip_route( + self.packet_plane.clone(), + self.global_ctx.clone(), + self.tun_dev.clone(), + self.tun_inet, + self.fake_ip, + ) + .await + } else { + MagicDnsServerInstance::new( + self.packet_plane.clone(), + self.global_ctx.clone(), + self.tun_dev.clone(), + self.tun_inet, + self.fake_ip, + ) + .await + }; + + match server_result { Ok(server) => { self.server = Some(server); tracing::info!("DnsRunner::run_once: server started"); @@ -74,11 +170,16 @@ impl DnsRunner { } pub async fn run(&mut self, canel_token: CancellationToken) { + let fake_ip_route_claim = self.fake_ip_route_claim(); + loop { tracing::info!("DnsRunner::run: start"); tokio::select! { _ = canel_token.cancelled() => { self.clean_env().await; + if let Some(claim) = &fake_ip_route_claim { + claim.remove().await; + } tracing::info!("DnsRunner::run: cancelled"); return; } diff --git a/easytier/src/instance/dns_server/server_instance.rs b/easytier/src/instance/dns_server/server_instance.rs index 4c1c3c00..2908f222 100644 --- a/easytier/src/instance/dns_server/server_instance.rs +++ b/easytier/src/instance/dns_server/server_instance.rs @@ -14,8 +14,10 @@ use super::{ }; use crate::{ common::{ + error::Error as EtError, global_ctx::ArcGlobalCtx, ifcfg::{IfConfiger, IfConfiguerTrait}, + netns::NetNS, }, instance::dns_server::{ config::{Record, RecordBuilder, RecordType}, @@ -51,7 +53,9 @@ use std::{collections::BTreeMap, io, net::Ipv4Addr, str::FromStr, sync::Arc, tim pub(super) struct MagicDnsServerInstanceData { dns_server: Server, tun_dev: Option, + net_ns: NetNS, fake_ip: Ipv4Addr, + manage_fake_ip_route: bool, route_store: MagicDnsRecordStore, record_apply: tokio::sync::Mutex<()>, @@ -356,12 +360,66 @@ fn get_system_config( } impl MagicDnsServerInstance { + async fn add_fake_ip_route( + tun_dev_name: &str, + fake_ip: Ipv4Addr, + net_ns: &NetNS, + cost: Option, + ) -> Result<(), anyhow::Error> { + let ifcfg = IfConfiger::default(); + let _guard = net_ns.guard(); + match ifcfg.add_ipv4_route(tun_dev_name, fake_ip, 32, cost).await { + Err(EtError::IOError(err)) if err.kind() == io::ErrorKind::AlreadyExists => { + ifcfg.remove_ipv4_route(tun_dev_name, fake_ip, 32).await?; + ifcfg + .add_ipv4_route(tun_dev_name, fake_ip, 32, cost) + .await?; + Ok(()) + } + ret => ret.map_err(Into::into), + } + } + + async fn remove_fake_ip_route(tun_dev_name: &str, fake_ip: Ipv4Addr, net_ns: &NetNS) { + let ifcfg = IfConfiger::default(); + let _guard = net_ns.guard(); + if let Err(err) = ifcfg.remove_ipv4_route(tun_dev_name, fake_ip, 32).await { + tracing::warn!( + ?err, + ?tun_dev_name, + ?fake_ip, + "remove magic dns route failed" + ); + } + } + pub(crate) async fn new( packet_plane: Arc, global_ctx: ArcGlobalCtx, tun_dev: Option, tun_inet: Ipv4Inet, fake_ip: Ipv4Addr, + ) -> Result { + Self::new_inner(packet_plane, global_ctx, tun_dev, tun_inet, fake_ip, true).await + } + + pub(crate) async fn new_with_external_fake_ip_route( + packet_plane: Arc, + global_ctx: ArcGlobalCtx, + tun_dev: Option, + tun_inet: Ipv4Inet, + fake_ip: Ipv4Addr, + ) -> Result { + Self::new_inner(packet_plane, global_ctx, tun_dev, tun_inet, fake_ip, false).await + } + + async fn new_inner( + packet_plane: Arc, + global_ctx: ArcGlobalCtx, + tun_dev: Option, + tun_inet: Ipv4Inet, + fake_ip: Ipv4Addr, + manage_fake_ip_route: bool, ) -> Result { let tcp_listener = runtime_rpc_listener(MAGIC_DNS_INSTANCE_SOCKET_ADDR.parse()?); let mut rpc_server = StandAloneServer::new(tcp_listener); @@ -374,7 +432,8 @@ impl MagicDnsServerInstance { let mut dns_server = Server::new(dns_config); dns_server.run().await?; - if !tun_inet.contains(&fake_ip) + if manage_fake_ip_route + && !tun_inet.contains(&fake_ip) && let Some(tun_dev_name) = &tun_dev { let cost = if cfg!(target_os = "windows") { @@ -382,16 +441,15 @@ impl MagicDnsServerInstance { } else { None }; - let ifcfg = IfConfiger {}; - ifcfg - .add_ipv4_route(tun_dev_name, fake_ip, 32, cost) - .await?; + Self::add_fake_ip_route(tun_dev_name, fake_ip, &global_ctx.net_ns, cost).await?; } let data = Arc::new(MagicDnsServerInstanceData { dns_server, tun_dev: tun_dev.clone(), + net_ns: global_ctx.net_ns.clone(), fake_ip, + manage_fake_ip_route, route_store: MagicDnsRecordStore::default(), record_apply: tokio::sync::Mutex::new(()), system_config: get_system_config(tun_dev.as_deref())?, @@ -436,14 +494,13 @@ impl MagicDnsServerInstance { if let Err(e) = ret { tracing::error!("Failed to close system config: {:?}", e); } - if !self.tun_inet.contains(&self.data.fake_ip) - && let Some(tun_dev_name) = &self.data.tun_dev - { - let ifcfg = IfConfiger {}; - let _ = ifcfg - .remove_ipv4_route(tun_dev_name, self.data.fake_ip, 32) - .await; - } + } + + if self.data.manage_fake_ip_route + && !self.tun_inet.contains(&self.data.fake_ip) + && let Some(tun_dev_name) = &self.data.tun_dev + { + Self::remove_fake_ip_route(tun_dev_name, self.data.fake_ip, &self.data.net_ns).await; } self.packet_filter.close().await; diff --git a/easytier/src/instance/factory.rs b/easytier/src/instance/factory.rs index 67dfb98c..1e546965 100644 --- a/easytier/src/instance/factory.rs +++ b/easytier/src/instance/factory.rs @@ -1,5 +1,10 @@ use std::sync::Arc; +#[cfg(feature = "tun")] +use tokio::sync::Mutex; + +#[cfg(all(feature = "management-rpc", feature = "tun", mobile))] +use easytier_core::instance::CoreInstanceState; #[cfg(any(feature = "management-rpc", test))] use easytier_core::instance::manager::InstanceManager; #[cfg(feature = "management-rpc")] @@ -12,6 +17,8 @@ use easytier_core::{ use crate::common::global_ctx::EventBusSubscriber; +#[cfg(feature = "tun")] +use super::shared_virtual_nic::{ArcSharedVirtualNicRegistry, SharedVirtualNicRegistry}; use super::{ composition::compose_native_core_instance, host::NativeInstanceHost, runtime_host::NativeInstanceRuntimeHost, @@ -51,6 +58,58 @@ pub fn subscribe_native_instance_event( .map(NativeInstanceRuntimeHost::subscribe_event) } +#[cfg(all(feature = "management-rpc", feature = "tun", mobile))] +pub async fn attach_mobile_tun_fd(manager: &NativeInstanceManager, fd: i32) -> anyhow::Result<()> { + let instances = manager + .instances() + .into_iter() + .filter(|instance| instance.state() == CoreInstanceState::Running) + .filter(|instance| { + instance + .runtime_host::() + .is_some_and(NativeInstanceRuntimeHost::tun_enabled) + }) + .collect::>(); + if fd > 0 && instances.is_empty() { + anyhow::bail!("no running TUN-enabled instance is available for fd attachment"); + } + + let mut errors = Vec::new(); + // The first member opens the TUN; the remaining members join its dispatcher. + for (index, instance) in instances.iter().enumerate() { + if let Err(error) = attach_mobile_tun_fd_to_instance(instance, fd, index == 0).await { + errors.push(format!("{}: {error}", instance.instance_id())); + if fd > 0 { + break; + } + } + } + if errors.is_empty() { + return Ok(()); + } + + if fd > 0 { + for instance in &instances { + if let Err(error) = attach_mobile_tun_fd_to_instance(instance, 0, false).await { + errors.push(format!("cleanup {}: {error}", instance.instance_id())); + } + } + } + anyhow::bail!("failed to attach mobile TUN fd: {}", errors.join("; ")) +} + +#[cfg(all(feature = "management-rpc", feature = "tun", mobile))] +async fn attach_mobile_tun_fd_to_instance( + instance: &NativeCoreInstance, + fd: i32, + replace_tun_fd: bool, +) -> anyhow::Result<()> { + let runtime = instance + .runtime_host::() + .ok_or_else(|| anyhow::anyhow!("native runtime host is unavailable"))?; + runtime.attach_mobile_tun_fd(fd, replace_tun_fd).await +} + #[cfg(feature = "management-rpc")] pub fn native_instance_manager_with_runtime( runtime_handle: tokio::runtime::Handle, @@ -97,6 +156,8 @@ fn native_instance_manager_with_optional_runtime( /// Native construction Adapter for the canonical core InstanceManager. pub struct NativeInstanceFactory { process_runtime: Arc, + #[cfg(feature = "tun")] + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, runtime_handle: Option, compact_runtime: bool, #[cfg(feature = "logging")] @@ -107,6 +168,8 @@ impl NativeInstanceFactory { pub fn new(process_runtime: Arc) -> Self { Self { process_runtime, + #[cfg(feature = "tun")] + shared_virtual_nic_registry: Arc::new(Mutex::new(SharedVirtualNicRegistry::new())), runtime_handle: None, compact_runtime: false, #[cfg(feature = "logging")] @@ -126,6 +189,7 @@ impl NativeInstanceFactory { self } + #[cfg(feature = "management-rpc")] fn with_compact_runtime(mut self) -> Self { self.compact_runtime = true; self @@ -149,6 +213,8 @@ impl InstanceFactory for NativeInstanceFactory { let instance = compose_native_core_instance( config, self.process_runtime.clone(), + #[cfg(feature = "tun")] + self.shared_virtual_nic_registry.clone(), self.compact_runtime, )?; #[cfg(feature = "logging")] diff --git a/easytier/src/instance/mod.rs b/easytier/src/instance/mod.rs index 39dc3200..bda8ca12 100644 --- a/easytier/src/instance/mod.rs +++ b/easytier/src/instance/mod.rs @@ -18,6 +18,9 @@ pub(crate) mod listeners; #[cfg(feature = "public-ipv6-provider")] pub(crate) mod public_ipv6_provider; +#[cfg(feature = "tun")] +pub mod shared_virtual_nic; + #[cfg(feature = "tun")] pub mod virtual_nic; diff --git a/easytier/src/instance/runtime_host.rs b/easytier/src/instance/runtime_host.rs index 93c0d7fc..ad0d80f4 100644 --- a/easytier/src/instance/runtime_host.rs +++ b/easytier/src/instance/runtime_host.rs @@ -1,13 +1,16 @@ use std::sync::Arc; +#[cfg(feature = "web-client")] +use easytier_core::config::runtime::CoreInstanceRuntimeConfig; use easytier_core::{ - config::runtime::CoreInstanceRuntimeConfig, gateway::dhcp::DhcpIpv4Host, - host::packet::HostPacketReceiver, instance::CorePacketPlane, + gateway::dhcp::DhcpIpv4Host, host::packet::HostPacketReceiver, instance::CorePacketPlane, }; use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; use crate::common::global_ctx::ArcGlobalCtx; +#[cfg(feature = "tun")] +use crate::instance::shared_virtual_nic::ArcSharedVirtualNicRegistry; mod event_journal; mod implementation; @@ -39,9 +42,17 @@ pub(crate) struct NativeInstanceRuntimeHost { } impl NativeInstanceRuntimeHost { - pub(crate) fn new(global_ctx: ArcGlobalCtx) -> Arc { + pub(crate) fn new( + global_ctx: ArcGlobalCtx, + #[cfg(feature = "tun")] shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, + ) -> Arc { let cancel = CancellationToken::new(); - let tun = NativeTunRuntime::new(global_ctx.clone(), cancel.clone()); + let tun = NativeTunRuntime::new( + global_ctx.clone(), + cancel.clone(), + #[cfg(feature = "tun")] + shared_virtual_nic_registry, + ); let event_journal = EventJournal::new(&global_ctx); Arc::new(Self { global_ctx, @@ -117,10 +128,24 @@ impl NativeInstanceRuntimeHost { self.global_ctx.subscribe() } + #[cfg(all(feature = "tun", mobile))] + pub(crate) fn tun_enabled(&self) -> bool { + !self.global_ctx.get_flags().no_tun + } + fn attach_runtime_tun_fd(&self, fd: i32) -> anyhow::Result<()> { self.tun.attach_fd(fd) } + #[cfg(all(feature = "tun", mobile))] + pub(crate) async fn attach_mobile_tun_fd( + &self, + fd: i32, + replace_tun_fd: bool, + ) -> anyhow::Result<()> { + self.tun.attach_mobile_fd(fd, replace_tun_fd).await + } + fn install_packet_receiver(&self, receiver: HostPacketReceiver) -> anyhow::Result<()> { self.tun.install_packet_receiver(receiver) } @@ -134,6 +159,16 @@ mod tests { global_ctx::{GlobalCtx, GlobalCtxEvent}, }; + fn runtime_host(global_ctx: ArcGlobalCtx) -> Arc { + NativeInstanceRuntimeHost::new( + global_ctx, + #[cfg(feature = "tun")] + Arc::new(tokio::sync::Mutex::new( + crate::instance::shared_virtual_nic::SharedVirtualNicRegistry::new(), + )), + ) + } + #[cfg(feature = "web-client")] fn runtime_config(config: &TomlConfig) -> CoreInstanceRuntimeConfig { let normalized = easytier_core::instance::CoreInstanceConfig::from_toml(config).unwrap(); @@ -146,7 +181,7 @@ mod tests { #[test] fn runtime_host_owns_event_subscription_context() { let global_ctx = Arc::new(GlobalCtx::new(TomlConfig::default())); - let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone()); + let runtime_host = runtime_host(global_ctx.clone()); let mut events = runtime_host.subscribe_event(); global_ctx.issue_event(GlobalCtxEvent::CredentialChanged); @@ -167,7 +202,7 @@ mod tests { config.set_ipv4(Some("10.20.0.1/24".parse().unwrap())); config.set_ipv6(Some("fd00::1/64".parse().unwrap())); let global_ctx = Arc::new(GlobalCtx::new(config.clone())); - let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone()); + let runtime_host = runtime_host(global_ctx.clone()); assert_eq!(global_ctx.get_hostname(), "before"); assert_eq!(global_ctx.get_ipv4(), Some("10.20.0.1/24".parse().unwrap())); @@ -209,7 +244,7 @@ mod tests { let config = TomlConfig::default(); config.set_dhcp(true); let global_ctx = Arc::new(GlobalCtx::new(config.clone())); - let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone()); + let runtime_host = runtime_host(global_ctx.clone()); let lease = "10.20.0.7/24".parse().unwrap(); global_ctx.set_ipv4(Some(lease)); diff --git a/easytier/src/instance/runtime_host/magic_dns.rs b/easytier/src/instance/runtime_host/magic_dns.rs index 8ee17048..0035ee2d 100644 --- a/easytier/src/instance/runtime_host/magic_dns.rs +++ b/easytier/src/instance/runtime_host/magic_dns.rs @@ -3,9 +3,9 @@ use easytier_core::instance::CorePacketPlane; #[cfg(feature = "magic-dns")] use tokio_util::{sync::CancellationToken, task::AbortOnDropHandle}; -use crate::common::global_ctx::ArcGlobalCtx; #[cfg(feature = "magic-dns")] use crate::instance::dns_server::{MAGIC_DNS_FAKE_IP, runner::DnsRunner}; +use crate::{common::global_ctx::ArcGlobalCtx, instance::virtual_nic::NicBackend}; #[derive(Default)] pub(super) struct MagicDnsRuntime { @@ -26,6 +26,7 @@ impl MagicDnsRuntime { packet_plane: std::sync::Arc, tun_dev: Option, tun_ip: Ipv4Inet, + shared_route_backend: Option, ) -> Self { let active = global_ctx.get_flags().accept_dns.then(|| { let mut runner = DnsRunner::new( @@ -34,7 +35,8 @@ impl MagicDnsRuntime { tun_dev, tun_ip, MAGIC_DNS_FAKE_IP.parse().unwrap(), - ); + ) + .with_shared_route_backend(shared_route_backend); let cancel = CancellationToken::new(); let task_cancel = cancel.clone(); let task = tokio::spawn(async move { @@ -54,6 +56,7 @@ impl MagicDnsRuntime { _packet_plane: std::sync::Arc, _tun_dev: Option, _tun_ip: Ipv4Inet, + _shared_route_backend: Option, ) -> Self { Self::default() } diff --git a/easytier/src/instance/runtime_host/tun_common.rs b/easytier/src/instance/runtime_host/tun_common.rs index a86aead9..44ac457f 100644 --- a/easytier/src/instance/runtime_host/tun_common.rs +++ b/easytier/src/instance/runtime_host/tun_common.rs @@ -3,11 +3,43 @@ use std::{ sync::{Arc, OnceLock}, }; -use easytier_core::host::packet::HostPacketReceiver; +use easytier_core::{host::packet::HostPacketReceiver, instance::CorePacketPlane}; use tokio::{sync::Mutex, task::JoinSet}; use super::MagicDnsRuntime; -use crate::instance::virtual_nic::NicCtx; +use crate::{ + common::{error::Error, global_ctx::ArcGlobalCtx}, + instance::{shared_virtual_nic::ArcSharedVirtualNicRegistry, virtual_nic::NicCtx}, +}; + +pub(super) async fn create_nic_ctx( + global_ctx: ArcGlobalCtx, + packet_plane: Arc, + receiver: Arc>, + close_notifier: Arc, + registry: ArcSharedVirtualNicRegistry, +) -> Result { + #[cfg(not(mobile))] + if global_ctx.get_flags().dev_name.is_empty() { + return Ok(NicCtx::new( + global_ctx, + packet_plane, + receiver, + close_notifier, + )); + } + + let member_id = global_ctx.get_id(); + NicCtx::new_shared( + global_ctx, + packet_plane, + receiver, + close_notifier, + registry, + member_id, + ) + .await +} struct NicCtxContainer { _nic_ctx: Option>, diff --git a/easytier/src/instance/runtime_host/tun_desktop.rs b/easytier/src/instance/runtime_host/tun_desktop.rs index 94629a31..f179750c 100644 --- a/easytier/src/instance/runtime_host/tun_desktop.rs +++ b/easytier/src/instance/runtime_host/tun_desktop.rs @@ -14,28 +14,37 @@ use tokio::{ }; use tokio_util::sync::CancellationToken; -use super::{MagicDnsRuntime, tun_common::TunNicState}; +use super::{ + MagicDnsRuntime, + tun_common::{TunNicState, create_nic_ctx}, +}; use crate::{ common::{ error::Error, global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, }, - instance::virtual_nic::NicCtx, + instance::shared_virtual_nic::ArcSharedVirtualNicRegistry, }; pub(super) struct NativeTunRuntime { global_ctx: ArcGlobalCtx, cancel: CancellationToken, nic: TunNicState, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, static_ip_task: Mutex>>, } impl NativeTunRuntime { - pub(super) fn new(global_ctx: ArcGlobalCtx, cancel: CancellationToken) -> Self { + pub(super) fn new( + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, + ) -> Self { Self { global_ctx, cancel, nic: TunNicState::empty(), + shared_virtual_nic_registry, static_ip_task: Mutex::new(None), } } @@ -67,6 +76,7 @@ impl NativeTunRuntime { let cancel = self.cancel.clone(); let global_ctx = self.global_ctx.clone(); let receiver = self.nic.receiver(); + let shared_virtual_nic_registry = self.shared_virtual_nic_registry.clone(); let (output, first_round) = oneshot::channel(); let task = tokio::spawn(async move { let mut output = Some(output); @@ -76,12 +86,29 @@ impl NativeTunRuntime { return; } let closed = Arc::new(Notify::new()); - let mut nic = NicCtx::new( + let mut nic = match create_nic_ctx( global_ctx.clone(), packet_plane.clone(), receiver.clone(), closed.clone(), - ); + shared_virtual_nic_registry.clone(), + ) + .await + { + Ok(nic) => nic, + Err(error) => { + if let Some(output) = output.take() { + let _ = output.send(Err(error)); + return; + } + tracing::error!(?error, "failed to create native interface context"); + tokio::select! { + _ = cancel.cancelled() => return, + _ = tokio::time::sleep(Duration::from_secs(1)) => {} + } + continue; + } + }; let result = tokio::select! { biased; _ = cancel.cancelled() => { @@ -104,11 +131,13 @@ impl NativeTunRuntime { } let magic_dns = if let Some(ip) = ipv4 { + let shared_route_backend = nic.shared_route_backend_for_dns(); MagicDnsRuntime::start( global_ctx.clone(), packet_plane.clone(), nic.ifname().await, ip, + shared_route_backend, ) } else { MagicDnsRuntime::default() @@ -161,6 +190,7 @@ impl NativeTunRuntime { nic: self.nic.clone(), closed: Arc::new(Notify::new()), packet_plane, + shared_virtual_nic_registry: self.shared_virtual_nic_registry.clone(), }) } } @@ -172,6 +202,7 @@ struct NativeDhcpIpv4Host { nic: TunNicState, closed: Arc, packet_plane: Arc, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, } impl NativeDhcpIpv4Host { @@ -199,21 +230,25 @@ impl NativeDhcpIpv4Host { return Ok(Some(ip)); } - let mut nic = NicCtx::new( + let mut nic = create_nic_ctx( self.global_ctx.clone(), self.packet_plane.clone(), self.nic.receiver(), self.closed.clone(), - ); + self.shared_virtual_nic_registry.clone(), + ) + .await?; tokio::select! { _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; DHCP update cancelled"), result = nic.run(Some(ip), self.global_ctx.get_ipv6()) => result?, } + let shared_route_backend = nic.shared_route_backend_for_dns(); let magic_dns = MagicDnsRuntime::start( self.global_ctx.clone(), self.packet_plane.clone(), nic.ifname().await, ip, + shared_route_backend, ); self.nic.install(nic, magic_dns).await; self.global_ctx.set_ipv4(Some(ip)); diff --git a/easytier/src/instance/runtime_host/tun_mobile.rs b/easytier/src/instance/runtime_host/tun_mobile.rs index 4f459326..900aad1f 100644 --- a/easytier/src/instance/runtime_host/tun_mobile.rs +++ b/easytier/src/instance/runtime_host/tun_mobile.rs @@ -8,26 +8,40 @@ use easytier_core::{ instance::CorePacketPlane, }; use futures::FutureExt as _; -use tokio::sync::{Mutex, Notify, mpsc}; +use tokio::sync::{Mutex, Notify, mpsc, oneshot}; use tokio_util::sync::CancellationToken; -use super::{MagicDnsRuntime, tun_common::TunNicState}; +use super::{ + MagicDnsRuntime, + tun_common::{TunNicState, create_nic_ctx}, +}; use crate::{ common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, - instance::virtual_nic::NicCtx, + instance::shared_virtual_nic::ArcSharedVirtualNicRegistry, }; +struct MobileTunAttachment { + fd: i32, + replace_tun_fd: bool, + completion: Option>>, +} + pub(super) struct NativeTunRuntime { global_ctx: ArcGlobalCtx, cancel: CancellationToken, nic: TunNicState, - tun_fd: mpsc::Sender, - tun_fd_receiver: Mutex>>, + tun_fd: mpsc::Sender, + tun_fd_receiver: Mutex>>, task: Mutex>>, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, } impl NativeTunRuntime { - pub(super) fn new(global_ctx: ArcGlobalCtx, cancel: CancellationToken) -> Self { + pub(super) fn new( + global_ctx: ArcGlobalCtx, + cancel: CancellationToken, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, + ) -> Self { let (tun_fd, tun_fd_receiver) = mpsc::channel(16); Self { global_ctx, @@ -36,6 +50,7 @@ impl NativeTunRuntime { tun_fd, tun_fd_receiver: Mutex::new(Some(tun_fd_receiver)), task: Mutex::new(None), + shared_virtual_nic_registry, } } @@ -51,23 +66,31 @@ impl NativeTunRuntime { global_ctx: ArcGlobalCtx, packet_plane: Arc, fd: i32, + replace_tun_fd: bool, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, ) -> anyhow::Result<()> { nic_state.drain().await; if fd <= 0 { return Ok(()); } let closed = Arc::new(Notify::new()); - let mut nic = NicCtx::new( + let mut nic = create_nic_ctx( global_ctx.clone(), packet_plane.clone(), nic_state.receiver(), closed, - ); - nic.run_for_mobile(fd).await.context("add ip failed")?; - let magic_dns = global_ctx - .get_ipv4() - .map(|ip| MagicDnsRuntime::start(global_ctx, packet_plane, None, ip)) - .unwrap_or_default(); + shared_virtual_nic_registry, + ) + .await?; + nic.run_for_mobile(fd, replace_tun_fd) + .await + .context("attach mobile TUN failed")?; + let magic_dns = if let Some(ip) = global_ctx.get_ipv4() { + let shared_route_backend = nic.shared_route_backend_for_dns(); + MagicDnsRuntime::start(global_ctx, packet_plane, None, ip, shared_route_backend) + } else { + MagicDnsRuntime::default() + }; nic_state.install(nic, magic_dns).await; Ok(()) } @@ -80,28 +103,43 @@ impl NativeTunRuntime { let nic_state = self.nic.clone(); let global_ctx = self.global_ctx.clone(); let cancel = self.cancel.clone(); + let shared_virtual_nic_registry = self.shared_virtual_nic_registry.clone(); self.task.lock().await.replace(tokio::spawn(async move { loop { - let fd = tokio::select! { + let attachment = tokio::select! { _ = cancel.cancelled() => return, - fd = tun_fds.recv() => match fd { Some(fd) => fd, None => return }, + attachment = tun_fds.recv() => match attachment { + Some(attachment) => attachment, + None => return, + }, }; - if let Err(error) = Self::install_mobile_tun( - nic_state.clone(), - global_ctx.clone(), - packet_plane.clone(), - fd, - ) - .await - { + let result = if attachment.fd <= 0 { + nic_state.drain().await; + Ok(()) + } else { + Self::install_mobile_tun( + nic_state.clone(), + global_ctx.clone(), + packet_plane.clone(), + attachment.fd, + attachment.replace_tun_fd, + shared_virtual_nic_registry.clone(), + ) + .await + }; + if let Err(error) = &result { tracing::error!(?error, "failed to attach mobile TUN fd"); } + if let Some(completion) = attachment.completion { + let _ = completion.send(result); + } } })); Ok(()) } pub(super) async fn shutdown(&self) { + self.cancel.cancel(); if let Some(task) = self.task.lock().await.take() { let _ = task.await; } @@ -110,10 +148,40 @@ impl NativeTunRuntime { pub(super) fn attach_fd(&self, fd: i32) -> anyhow::Result<()> { self.tun_fd - .try_send(fd) + .try_send(MobileTunAttachment { + fd, + replace_tun_fd: true, + completion: None, + }) .map_err(|error| anyhow::anyhow!("failed to send TUN fd: {error}")) } + pub(super) async fn attach_mobile_fd( + &self, + fd: i32, + replace_tun_fd: bool, + ) -> anyhow::Result<()> { + if self.task.lock().await.is_none() { + anyhow::bail!("mobile TUN runtime is not running"); + } + + let (completion, result) = oneshot::channel(); + tokio::select! { + _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; TUN attachment cancelled"), + send_result = self.tun_fd.send(MobileTunAttachment { + fd, + replace_tun_fd, + completion: Some(completion), + }) => send_result.map_err(|error| anyhow::anyhow!("failed to send TUN fd: {error}"))?, + } + + tokio::select! { + _ = self.cancel.cancelled() => anyhow::bail!("instance is closing; TUN attachment cancelled"), + result = result => result + .map_err(|_| anyhow::anyhow!("mobile TUN runtime stopped before attachment completed"))?, + } + } + pub(super) fn dhcp_host( &self, operation: Arc>, diff --git a/easytier/src/instance/shared_virtual_nic.rs b/easytier/src/instance/shared_virtual_nic.rs new file mode 100644 index 00000000..2d643dc0 --- /dev/null +++ b/easytier/src/instance/shared_virtual_nic.rs @@ -0,0 +1,1738 @@ +use std::{ + collections::{BTreeMap, BTreeSet}, + net::{Ipv4Addr, Ipv6Addr}, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, +}; + +use cidr::{Ipv4Cidr, Ipv4Inet, Ipv6Cidr, Ipv6Inet}; +use tokio::sync::{Mutex, Notify}; + +use crate::common::error::Error; +use easytier_core::tunnel::{Tunnel, ring::create_ring_tunnel_pair}; + +use super::virtual_nic::{VirtualNic, VirtualNicConfig}; + +#[cfg(not(target_os = "linux"))] +use crate::common::ifcfg::IfConfiger; + +mod dispatcher; + +use dispatcher::{SharedVirtualNicDispatcher, SharedVirtualNicMemberTunnelTable}; + +pub type SharedVirtualNicMemberId = uuid::Uuid; +pub(super) type SharedVirtualNicMemberRegistrationId = uuid::Uuid; +pub(crate) type ArcSharedVirtualNicRegistry = Arc>; + +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub struct SharedIpv4Route { + pub address: Ipv4Addr, + pub prefix: u8, + pub cost: Option, +} + +impl SharedIpv4Route { + pub fn new(address: Ipv4Addr, prefix: u8, cost: Option) -> Self { + Self { + address, + prefix, + cost, + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub struct SharedIpv6Route { + pub address: Ipv6Addr, + pub prefix: u8, + pub cost: Option, +} + +impl SharedIpv6Route { + pub fn new(address: Ipv6Addr, prefix: u8, cost: Option) -> Self { + Self { + address, + prefix, + cost, + } + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct SharedIfConfigClaims { + pub ipv4_addresses: BTreeSet, + pub ipv6_addresses: BTreeSet, + pub ipv4_routes: BTreeSet, + pub ipv6_routes: BTreeSet, + pub mtu: Option, +} + +impl SharedIfConfigClaims { + fn ipv4_destinations(&self) -> impl Iterator + '_ { + self.ipv4_addresses.iter().map(Ipv4Inet::network).chain( + self.ipv4_routes + .iter() + .filter_map(|route| Ipv4Inet::new(route.address, route.prefix).ok()) + .map(|route| route.network()), + ) + } + + fn ipv6_destinations(&self) -> impl Iterator + '_ { + self.ipv6_addresses.iter().map(Ipv6Inet::network).chain( + self.ipv6_routes + .iter() + .filter_map(|route| Ipv6Inet::new(route.address, route.prefix).ok()) + .map(|route| route.network()), + ) + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct OwnedItemDelta { + pub added: BTreeSet, + pub removed: BTreeSet, +} + +impl Default for OwnedItemDelta { + fn default() -> Self { + Self { + added: BTreeSet::new(), + removed: BTreeSet::new(), + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct SharedMtuChange { + pub old: Option, + pub new: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct SharedIfConfigDelta { + pub ipv4_addresses: OwnedItemDelta, + pub ipv6_addresses: OwnedItemDelta, + pub ipv4_routes: OwnedItemDelta, + pub ipv6_routes: OwnedItemDelta, + pub mtu: Option, +} + +impl SharedIfConfigDelta { + fn between(old: &EffectiveIfConfig, new: &EffectiveIfConfig) -> Self { + Self { + ipv4_addresses: item_delta(&old.ipv4_addresses, &new.ipv4_addresses), + ipv6_addresses: item_delta(&old.ipv6_addresses, &new.ipv6_addresses), + ipv4_routes: item_delta(&old.ipv4_routes, &new.ipv4_routes), + ipv6_routes: item_delta(&old.ipv6_routes, &new.ipv6_routes), + mtu: mtu_delta(old.effective_mtu, new.effective_mtu), + } + } +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +struct EffectiveIfConfig { + ipv4_addresses: BTreeSet, + ipv6_addresses: BTreeSet, + ipv4_routes: BTreeSet, + ipv6_routes: BTreeSet, + effective_mtu: Option, +} + +#[derive(Clone, Debug, Default)] +pub struct SharedIfConfig { + member_claims: BTreeMap, +} + +impl SharedIfConfig { + fn ensure_disjoint( + &self, + member_id: SharedVirtualNicMemberId, + claims: &SharedIfConfigClaims, + ) -> anyhow::Result<()> { + let magic_dns = Ipv4Cidr::new(crate::instance::dns_server::MAGIC_DNS_FAKE_IP.parse()?, 32)?; + for (other_id, other) in &self.member_claims { + if *other_id == member_id { + continue; + } + for requested in claims.ipv4_destinations() { + if let Some(existing) = other.ipv4_destinations().find(|existing| { + !(requested == magic_dns && *existing == magic_dns) + && ipv4_cidrs_overlap(requested, *existing) + }) { + anyhow::bail!( + "shared virtual NIC destination {requested} overlaps member {other_id} destination {existing}" + ); + } + } + for requested in claims.ipv6_destinations() { + if let Some(existing) = other + .ipv6_destinations() + .find(|existing| ipv6_cidrs_overlap(requested, *existing)) + { + anyhow::bail!( + "shared virtual NIC destination {requested} overlaps member {other_id} destination {existing}" + ); + } + } + } + Ok(()) + } + + fn apply_member_claims( + &mut self, + member_id: SharedVirtualNicMemberId, + claims: SharedIfConfigClaims, + ) -> SharedIfConfigDelta { + let old = self.effective(); + self.member_claims.insert(member_id, claims); + SharedIfConfigDelta::between(&old, &self.effective()) + } + + pub fn remove_member( + &mut self, + member_id: SharedVirtualNicMemberId, + ) -> Option { + let old = self.effective(); + self.member_claims.remove(&member_id)?; + Some(SharedIfConfigDelta::between(&old, &self.effective())) + } + + pub fn effective_mtu(&self) -> Option { + self.member_claims + .values() + .filter_map(|claims| claims.mtu) + .min() + } + + pub fn owners_of_ipv4_route( + &self, + route: &SharedIpv4Route, + ) -> BTreeSet { + owners_for(&self.member_claims, |claims| &claims.ipv4_routes, route) + } + + pub fn owners_of_ipv4_address(&self, address: &Ipv4Inet) -> BTreeSet { + owners_for( + &self.member_claims, + |claims| &claims.ipv4_addresses, + address, + ) + } + + fn effective(&self) -> EffectiveIfConfig { + EffectiveIfConfig { + ipv4_addresses: self + .member_claims + .values() + .flat_map(|claims| claims.ipv4_addresses.iter().copied()) + .collect(), + ipv6_addresses: self + .member_claims + .values() + .flat_map(|claims| claims.ipv6_addresses.iter().copied()) + .collect(), + ipv4_routes: self + .member_claims + .values() + .flat_map(|claims| claims.ipv4_routes.iter().cloned()) + .collect(), + ipv6_routes: self + .member_claims + .values() + .flat_map(|claims| claims.ipv6_routes.iter().cloned()) + .collect(), + effective_mtu: self.effective_mtu(), + } + } + + fn claims_of(&self, member_id: SharedVirtualNicMemberId) -> SharedIfConfigClaims { + self.member_claims + .get(&member_id) + .cloned() + .unwrap_or_default() + } + + fn ipv4_route_source_hint(&self, route: &SharedIpv4Route) -> Option { + let route_inet = Ipv4Inet::new(route.address, route.prefix).ok(); + let mut fallback = None; + + for claims in self.member_claims.values() { + if !claims.ipv4_routes.contains(route) { + continue; + } + + for address in &claims.ipv4_addresses { + fallback.get_or_insert(address.address()); + if route_inet + .as_ref() + .is_some_and(|route_inet| route_inet.contains(&address.address())) + { + return Some(address.address()); + } + } + } + + fallback + } + + fn changed_ipv4_route_sources(&self, next: &Self) -> BTreeSet { + let old_routes = self.effective().ipv4_routes; + let next_routes = next.effective().ipv4_routes; + old_routes + .iter() + .filter(|route| { + next_routes.contains(*route) + && self.ipv4_route_source_hint(route) != next.ipv4_route_source_hint(route) + }) + .cloned() + .collect() + } +} + +pub struct SharedVirtualNic { + nic: Arc>, + ifcfg: SharedIfConfig, + valid: Arc, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + member_registrations: BTreeMap, + dispatcher: Option, +} + +impl SharedVirtualNic { + pub fn new(config: VirtualNicConfig) -> Self { + Self { + nic: Arc::new(Mutex::new(VirtualNic::new(config))), + ifcfg: SharedIfConfig::default(), + valid: Arc::new(AtomicBool::new(true)), + member_tunnel_table: SharedVirtualNicMemberTunnelTable::default(), + member_registrations: BTreeMap::new(), + dispatcher: None, + } + } + + pub fn mark_invalid(&self) { + self.valid.store(false, Ordering::Release); + } + + pub fn is_valid(&self) -> bool { + self.valid.load(Ordering::Acquire) + } + + pub fn ifcfg(&self) -> &SharedIfConfig { + &self.ifcfg + } + + pub fn ifcfg_mut(&mut self) -> &mut SharedIfConfig { + &mut self.ifcfg + } + + pub fn nic(&self) -> Arc> { + self.nic.clone() + } + + #[cfg(not(target_os = "linux"))] + async fn ifcfg_and_ifname(&self) -> Result<(IfConfiger, String), Error> { + self.ensure_valid()?; + let nic = self.nic.lock().await; + Ok((nic.get_ifcfg(), nic.ifname().to_owned())) + } + + async fn link_up(&self) -> Result<(), Error> { + self.ensure_valid()?; + self.nic.lock().await.link_up().await + } + + async fn attach_member_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> Result<(), Error> { + self.ensure_valid()?; + + match self.member_registrations.get(&member_id).copied() { + Some(old_registration_id) if old_registration_id == registration_id => { + return Ok(()); + } + Some(_) => { + if let Err(err) = self.remove_member_claims(member_id).await { + self.invalidate_and_shutdown_dispatcher().await; + return Err(err); + } + } + None => {} + } + + self.member_registrations.insert(member_id, registration_id); + Ok(()) + } + + fn is_current_member_registration( + &self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> bool { + self.member_registrations + .get(&member_id) + .is_some_and(|current| *current == registration_id) + } + + async fn apply_member_claims_for_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + self.apply_member_claims(member_id, claims).await + } + + async fn remove_member_registration_claims( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + + if let Err(err) = self.remove_member_claims(member_id).await { + self.member_registrations.remove(&member_id); + self.invalidate_and_shutdown_dispatcher().await; + return Err(err); + } + + self.member_registrations.remove(&member_id); + self.shutdown_dispatcher_if_idle().await; + Ok(()) + } + + async fn apply_member_claims( + &mut self, + member_id: SharedVirtualNicMemberId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + self.ensure_valid()?; + self.ifcfg.ensure_disjoint(member_id, &claims)?; + + let mut next_ifcfg = self.ifcfg.clone(); + let old_claims = self.ifcfg.claims_of(member_id); + let next_claims = claims.clone(); + let delta = next_ifcfg.apply_member_claims(member_id, claims); + self.sync_dispatcher_sources_for_ifcfg_update(member_id, &old_claims, &next_claims) + .await?; + + if let Err(err) = self.apply_ifcfg_delta(&delta, &next_ifcfg).await { + let _ = self + .sync_dispatcher_sources_for_claims(member_id, &old_claims) + .await; + return Err(err); + } + + self.sync_dispatcher_sources_for_claims(member_id, &next_claims) + .await?; + self.ifcfg = next_ifcfg; + + Ok(()) + } + + #[cfg(mobile)] + async fn apply_member_claims_for_mobile( + &mut self, + member_id: SharedVirtualNicMemberId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + self.ensure_valid()?; + self.ifcfg.ensure_disjoint(member_id, &claims)?; + + let mut next_ifcfg = self.ifcfg.clone(); + let next_claims = claims.clone(); + next_ifcfg.apply_member_claims(member_id, claims); + + self.sync_dispatcher_sources_for_claims(member_id, &next_claims) + .await?; + self.ifcfg = next_ifcfg; + + Ok(()) + } + + async fn apply_member_mtu_for_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + mtu: u32, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + + let mut claims = self.ifcfg.claims_of(member_id); + claims.mtu = Some(mtu); + self.apply_member_claims(member_id, claims).await + } + + #[cfg(mobile)] + async fn apply_member_mtu_for_mobile_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + mtu: u32, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + + let mut claims = self.ifcfg.claims_of(member_id); + claims.mtu = Some(mtu); + self.apply_member_claims_for_mobile(member_id, claims).await + } + + async fn remove_member_claims( + &mut self, + member_id: SharedVirtualNicMemberId, + ) -> Result<(), Error> { + self.ensure_valid()?; + + let mut next_ifcfg = self.ifcfg.clone(); + let Some(delta) = next_ifcfg.remove_member(member_id) else { + return Ok(()); + }; + #[cfg(not(mobile))] + self.apply_ifcfg_delta(&delta, &next_ifcfg).await?; + #[cfg(mobile)] + drop(delta); + if let Some(dispatcher) = &self.dispatcher { + dispatcher.remove_sources(member_id).await?; + } + self.ifcfg = next_ifcfg; + + Ok(()) + } + + async fn shutdown_dispatcher_if_idle(&mut self) { + if !self.member_registrations.is_empty() { + return; + } + + self.shutdown_dispatcher().await; + } + + async fn invalidate_and_shutdown_dispatcher(&mut self) { + self.mark_invalid(); + self.shutdown_dispatcher().await; + } + + async fn shutdown_dispatcher(&mut self) { + if let Some(dispatcher) = self.dispatcher.take() { + dispatcher.shutdown_without_invalidation().await; + } + } + + async fn apply_ifcfg_delta( + &self, + delta: &SharedIfConfigDelta, + next_ifcfg: &SharedIfConfig, + ) -> Result<(), Error> { + let changed_routes = self.ifcfg.changed_ipv4_route_sources(next_ifcfg); + #[cfg(target_os = "linux")] + let next_effective = next_ifcfg.effective(); + let nic = self.nic.lock().await; + + for route in &delta.ipv4_routes.removed { + let source_hint = self.ifcfg.ipv4_route_source_hint(route); + ignore_removed_ifcfg_not_found( + remove_shared_ipv4_route(&nic, route, source_hint).await, + )?; + } + for route in &changed_routes { + let source_hint = self.ifcfg.ipv4_route_source_hint(route); + ignore_removed_ifcfg_not_found( + remove_shared_ipv4_route(&nic, route, source_hint).await, + )?; + } + for route in &delta.ipv6_routes.removed { + ignore_removed_ifcfg_not_found( + nic.remove_ipv6_route(route.address, route.prefix).await, + )?; + } + for ip in &delta.ipv4_addresses.removed { + ignore_removed_ifcfg_not_found(nic.remove_ip(Some(*ip)).await)?; + } + for ip in &delta.ipv6_addresses.removed { + ignore_removed_ifcfg_not_found(nic.remove_ipv6(Some(*ip)).await)?; + } + + for ip in &delta.ipv4_addresses.added { + nic.add_ip(ip.address(), ip.network_length() as i32).await?; + } + for ip in &delta.ipv6_addresses.added { + nic.add_ipv6(ip.address(), ip.network_length() as i32) + .await?; + } + for route in &delta.ipv4_routes.added { + add_shared_ipv4_route(&nic, route, next_ifcfg).await?; + } + for route in &changed_routes { + add_shared_ipv4_route(&nic, route, next_ifcfg).await?; + } + for route in &delta.ipv6_routes.added { + nic.add_ipv6_route_with_cost(route.address, route.prefix, route.cost) + .await?; + } + + if let Some(mtu) = &delta.mtu { + nic.set_mtu(mtu.new.unwrap_or_else(|| nic.configured_mtu())) + .await?; + } + + #[cfg(target_os = "linux")] + if !delta.ipv4_addresses.removed.is_empty() { + for route in &next_effective.ipv4_routes { + ignore_added_ifcfg_already_exists( + add_shared_ipv4_route(&nic, route, next_ifcfg).await, + )?; + } + } + + #[cfg(target_os = "linux")] + if !delta.ipv6_addresses.removed.is_empty() { + for route in &next_effective.ipv6_routes { + ignore_added_ifcfg_already_exists( + nic.add_ipv6_route_with_cost(route.address, route.prefix, route.cost) + .await, + )?; + } + } + + Ok(()) + } + + fn ensure_valid(&self) -> Result<(), Error> { + if self.is_valid() { + return Ok(()); + } + + Err(anyhow::anyhow!("shared virtual nic is invalid").into()) + } + + fn member_tunnel_table(&self) -> SharedVirtualNicMemberTunnelTable { + self.member_tunnel_table.clone() + } + + fn valid_flag(&self) -> Arc { + self.valid.clone() + } + + async fn ensure_dispatcher(&mut self) -> Result<(), Error> { + self.ensure_valid()?; + + if self.dispatcher.is_some() { + return Ok(()); + } + + let tunnel = self.nic.lock().await.create_dev().await?; + let dispatcher = SharedVirtualNicDispatcher::start( + tunnel, + self.member_tunnel_table.clone(), + self.valid.clone(), + ); + self.sync_dispatcher_sources(&dispatcher).await?; + self.dispatcher = Some(dispatcher); + Ok(()) + } + + #[cfg(mobile)] + async fn ensure_dispatcher_for_mobile( + &mut self, + tun_fd: std::os::fd::RawFd, + replace_tun_fd: bool, + ) -> Result<(), Error> { + self.ensure_valid()?; + + if let Some(dispatcher) = &self.dispatcher { + if replace_tun_fd { + dispatcher.update_mobile_tun_fd(tun_fd).await?; + } + return Ok(()); + } + + let dispatcher = SharedVirtualNicDispatcher::start_for_mobile( + self.nic.clone(), + tun_fd, + self.member_tunnel_table.clone(), + self.valid.clone(), + ) + .await?; + self.sync_dispatcher_sources(&dispatcher).await?; + self.dispatcher = Some(dispatcher); + Ok(()) + } + + async fn sync_dispatcher_sources( + &self, + dispatcher: &SharedVirtualNicDispatcher, + ) -> Result<(), Error> { + for (member_id, claims) in &self.ifcfg.member_claims { + dispatcher.update_sources(*member_id, claims).await?; + } + Ok(()) + } + + async fn sync_dispatcher_sources_for_ifcfg_update( + &self, + member_id: SharedVirtualNicMemberId, + old_claims: &SharedIfConfigClaims, + next_claims: &SharedIfConfigClaims, + ) -> Result<(), Error> { + let active_claims = dispatcher_claims_for_ifcfg_transition(old_claims, next_claims); + self.sync_dispatcher_sources_for_claims(member_id, &active_claims) + .await + } + + async fn sync_dispatcher_sources_for_claims( + &self, + member_id: SharedVirtualNicMemberId, + claims: &SharedIfConfigClaims, + ) -> Result<(), Error> { + if let Some(dispatcher) = &self.dispatcher { + dispatcher.update_sources(member_id, claims).await?; + } + Ok(()) + } +} + +fn dispatcher_claims_for_ifcfg_transition( + old_claims: &SharedIfConfigClaims, + next_claims: &SharedIfConfigClaims, +) -> SharedIfConfigClaims { + SharedIfConfigClaims { + ipv4_addresses: merged_items(&old_claims.ipv4_addresses, &next_claims.ipv4_addresses), + ipv6_addresses: merged_items(&old_claims.ipv6_addresses, &next_claims.ipv6_addresses), + ipv4_routes: merged_items(&old_claims.ipv4_routes, &next_claims.ipv4_routes), + ipv6_routes: merged_items(&old_claims.ipv6_routes, &next_claims.ipv6_routes), + mtu: None, + } +} + +async fn add_shared_ipv4_route( + nic: &VirtualNic, + route: &SharedIpv4Route, + ifcfg: &SharedIfConfig, +) -> Result<(), Error> { + nic.add_route_with_cost_and_source_hint( + route.address, + route.prefix, + route.cost, + ifcfg.ipv4_route_source_hint(route), + ) + .await +} + +async fn remove_shared_ipv4_route( + nic: &VirtualNic, + route: &SharedIpv4Route, + source_hint: Option, +) -> Result<(), Error> { + nic.remove_route_with_cost_and_source_hint(route.address, route.prefix, route.cost, source_hint) + .await +} + +fn merged_items(old_items: &BTreeSet, new_items: &BTreeSet) -> BTreeSet +where + T: Ord + Clone, +{ + old_items.union(new_items).cloned().collect() +} + +fn remove_claimed_item(items: &mut BTreeSet, item: Option) +where + T: Ord, +{ + match item { + Some(item) => { + items.remove(&item); + } + None => { + items.clear(); + } + } +} + +fn ignore_removed_ifcfg_not_found(result: Result<(), Error>) -> Result<(), Error> { + match result { + Err(Error::NotFound) => Ok(()), + other => other, + } +} + +#[cfg(target_os = "linux")] +fn ignore_added_ifcfg_already_exists(result: Result<(), Error>) -> Result<(), Error> { + match result { + Err(Error::IOError(err)) if err.kind() == std::io::ErrorKind::AlreadyExists => Ok(()), + other => other, + } +} + +struct SharedVirtualNicMemberRegistration { + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + shared_nic: Arc>, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, +} + +impl SharedVirtualNicMemberRegistration { + fn register_tunnel( + &self, + tunnel: Box, + close_notifier: Arc, + ) -> Result<(), Error> { + self.member_tunnel_table.register( + self.member_id, + self.registration_id, + tunnel, + close_notifier, + ) + } +} + +impl Drop for SharedVirtualNicMemberRegistration { + fn drop(&mut self) { + self.member_tunnel_table + .unregister(self.member_id, self.registration_id); + let shared_nic = self.shared_nic.clone(); + let member_id = self.member_id; + let registration_id = self.registration_id; + + let Ok(handle) = tokio::runtime::Handle::try_current() else { + tracing::warn!( + ?member_id, + "skip shared virtual nic member claim cleanup without tokio runtime" + ); + return; + }; + + handle.spawn(async move { + let mut shared_nic = shared_nic.lock().await; + if let Err(err) = shared_nic + .remove_member_registration_claims(member_id, registration_id) + .await + { + tracing::warn!( + ?member_id, + ?err, + "failed to clean shared virtual nic member claims" + ); + } + }); + } +} + +#[derive(Clone)] +pub struct SharedVirtualNicMember { + member_id: SharedVirtualNicMemberId, + configured_mtu: u32, + shared_nic: Arc>, + close_notifier: Arc, + registration: Arc, +} + +impl SharedVirtualNicMember { + fn new( + member_id: SharedVirtualNicMemberId, + configured_mtu: u32, + shared_nic: Arc>, + close_notifier: Arc, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + ) -> Self { + let registration_id = uuid::Uuid::new_v4(); + Self { + member_id, + configured_mtu, + shared_nic: shared_nic.clone(), + close_notifier, + registration: Arc::new(SharedVirtualNicMemberRegistration { + member_id, + registration_id, + shared_nic: shared_nic.clone(), + member_tunnel_table, + }), + } + } + + pub fn member_id(&self) -> SharedVirtualNicMemberId { + self.member_id + } + + pub fn shared_nic(&self) -> Arc> { + self.shared_nic.clone() + } + + pub fn close_notifier(&self) -> Arc { + self.close_notifier.clone() + } + + #[cfg(test)] + fn configured_mtu_for_test(&self) -> u32 { + self.configured_mtu + } + + pub async fn create_dev(&self) -> Result, Error> { + let (member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); + { + let mut shared_nic = self.shared_nic.lock().await; + shared_nic + .attach_member_registration(self.member_id, self.registration.registration_id) + .await?; + shared_nic.ensure_dispatcher().await?; + shared_nic + .apply_member_mtu_for_registration( + self.member_id, + self.registration.registration_id, + self.configured_mtu, + ) + .await?; + } + self.registration + .register_tunnel(shared_tunnel, self.close_notifier.clone())?; + Ok(member_tunnel) + } + + #[cfg(mobile)] + pub async fn create_dev_for_mobile( + &self, + tun_fd: std::os::fd::RawFd, + replace_tun_fd: bool, + ) -> Result, Error> { + let (member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); + { + let mut shared_nic = self.shared_nic.lock().await; + shared_nic + .attach_member_registration(self.member_id, self.registration.registration_id) + .await?; + shared_nic + .ensure_dispatcher_for_mobile(tun_fd, replace_tun_fd) + .await?; + shared_nic + .apply_member_mtu_for_mobile_registration( + self.member_id, + self.registration.registration_id, + self.configured_mtu, + ) + .await?; + } + self.registration + .register_tunnel(shared_tunnel, self.close_notifier.clone())?; + Ok(member_tunnel) + } + + #[cfg(not(target_os = "linux"))] + pub async fn ifcfg_and_ifname(&self) -> Result<(IfConfiger, String), Error> { + self.shared_nic.lock().await.ifcfg_and_ifname().await + } + + pub async fn link_up(&self) -> Result<(), Error> { + self.shared_nic.lock().await.link_up().await + } + + pub async fn add_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { + let ip = ipv4_inet(ip, cidr)?; + self.update_claims(|claims| { + claims.ipv4_addresses.insert(ip); + }) + .await + } + + pub async fn remove_ip(&self, ip: Option) -> Result<(), Error> { + self.update_claims(|claims| { + remove_claimed_item(&mut claims.ipv4_addresses, ip); + }) + .await + } + + pub async fn add_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { + let ip = ipv6_inet(ip, cidr)?; + self.update_claims(|claims| { + claims.ipv6_addresses.insert(ip); + }) + .await + } + + pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { + self.update_claims(|claims| { + remove_claimed_item(&mut claims.ipv6_addresses, ip); + }) + .await + } + + pub async fn add_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { + self.add_route_with_cost(address, cidr, None).await + } + + pub async fn add_route_with_cost( + &self, + address: Ipv4Addr, + cidr: u8, + cost: Option, + ) -> Result<(), Error> { + self.update_claims(|claims| { + claims + .ipv4_routes + .insert(SharedIpv4Route::new(address, cidr, cost)); + }) + .await + } + + pub async fn remove_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { + self.update_claims(|claims| { + claims + .ipv4_routes + .retain(|route| route.address != address || route.prefix != cidr); + }) + .await + } + + pub async fn add_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { + self.add_ipv6_route_with_cost(address, cidr, None).await + } + + pub async fn add_ipv6_route_with_cost( + &self, + address: Ipv6Addr, + cidr: u8, + cost: Option, + ) -> Result<(), Error> { + self.update_claims(|claims| { + claims + .ipv6_routes + .insert(SharedIpv6Route::new(address, cidr, cost)); + }) + .await + } + + pub async fn remove_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { + self.update_claims(|claims| { + claims + .ipv6_routes + .retain(|route| route.address != address || route.prefix != cidr); + }) + .await + } + + async fn update_claims(&self, update: F) -> Result<(), Error> + where + F: FnOnce(&mut SharedIfConfigClaims) + Send, + { + let mut shared_nic = self.shared_nic.lock().await; + let mut claims = shared_nic.ifcfg.claims_of(self.member_id); + update(&mut claims); + shared_nic + .apply_member_claims_for_registration( + self.member_id, + self.registration.registration_id, + claims, + ) + .await + } +} + +#[derive(Default)] +pub struct SharedVirtualNicRegistry { + nics: BTreeMap, +} + +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] +struct SharedVirtualNicRegistryKey { + net_ns: Option, + dev_name: String, +} + +impl SharedVirtualNicRegistryKey { + fn new(dev_name: String, config: &VirtualNicConfig) -> Self { + Self { + net_ns: config.net_ns_name(), + dev_name, + } + } +} + +struct SharedVirtualNicRegistryEntry { + nic: Arc>, + valid: Arc, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, +} + +impl SharedVirtualNicRegistryEntry { + fn new(nic: SharedVirtualNic) -> Self { + Self { + valid: nic.valid_flag(), + member_tunnel_table: nic.member_tunnel_table(), + nic: Arc::new(Mutex::new(nic)), + } + } + + fn is_valid(&self) -> bool { + self.valid.load(Ordering::Acquire) + } + + fn nic(&self) -> Arc> { + self.nic.clone() + } + + fn member_tunnel_table(&self) -> SharedVirtualNicMemberTunnelTable { + self.member_tunnel_table.clone() + } +} + +impl SharedVirtualNicRegistry { + pub fn new() -> Self { + Self::default() + } + + pub fn get( + &self, + dev_name: &str, + config: &VirtualNicConfig, + ) -> Option>> { + let key = SharedVirtualNicRegistryKey::new(dev_name.to_owned(), config); + self.nics + .get(&key) + .filter(|entry| entry.is_valid()) + .map(|entry| entry.nic()) + } + + #[cfg(test)] + pub fn get_by_dev_name_for_test(&self, dev_name: &str) -> Option>> { + let mut matches = self + .nics + .iter() + .filter(|(key, entry)| key.dev_name == dev_name && entry.is_valid()) + .map(|(_, entry)| entry.nic()); + let first = matches.next()?; + if matches.next().is_some() { + return None; + } + Some(first) + } + + pub fn get_or_create( + &mut self, + dev_name: String, + config: VirtualNicConfig, + ) -> Arc> { + self.get_or_create_entry(dev_name, config).nic() + } + + fn get_or_create_entry( + &mut self, + dev_name: String, + config: VirtualNicConfig, + ) -> &SharedVirtualNicRegistryEntry { + let key = SharedVirtualNicRegistryKey::new(dev_name, &config); + let needs_new_entry = self.nics.get(&key).is_none_or(|entry| !entry.is_valid()); + if needs_new_entry { + let entry = SharedVirtualNicRegistryEntry::new(SharedVirtualNic::new(config)); + self.nics.insert(key.clone(), entry); + } + + self.nics + .get(&key) + .expect("shared virtual nic registry entry should exist") + } + + pub fn create_member( + &mut self, + dev_name: String, + config: VirtualNicConfig, + member_id: SharedVirtualNicMemberId, + close_notifier: Arc, + ) -> SharedVirtualNicMember { + let configured_mtu = config.mtu(); + let entry = self.get_or_create_entry(dev_name, config); + SharedVirtualNicMember::new( + member_id, + configured_mtu, + entry.nic(), + close_notifier, + entry.member_tunnel_table(), + ) + } +} + +fn item_delta(old: &BTreeSet, new: &BTreeSet) -> OwnedItemDelta +where + T: Ord + Clone, +{ + OwnedItemDelta { + added: new.difference(old).cloned().collect(), + removed: old.difference(new).cloned().collect(), + } +} + +fn owners_for( + claims: &BTreeMap, + select: fn(&SharedIfConfigClaims) -> &BTreeSet, + item: &T, +) -> BTreeSet +where + T: Ord, +{ + claims + .iter() + .filter_map(|(member_id, claims)| select(claims).contains(item).then_some(*member_id)) + .collect() +} + +fn mtu_delta(old: Option, new: Option) -> Option { + (old != new).then_some(SharedMtuChange { old, new }) +} + +fn ipv4_cidrs_overlap(left: Ipv4Cidr, right: Ipv4Cidr) -> bool { + left.contains(&right.first_address()) || right.contains(&left.first_address()) +} + +fn ipv6_cidrs_overlap(left: Ipv6Cidr, right: Ipv6Cidr) -> bool { + left.contains(&right.first_address()) || right.contains(&left.first_address()) +} + +fn ipv4_inet(address: Ipv4Addr, prefix: i32) -> Result { + let prefix = u8::try_from(prefix) + .map_err(|_| anyhow::anyhow!("invalid IPv4 prefix length {}", prefix))?; + Ipv4Inet::new(address, prefix).map_err(|err| { + anyhow::anyhow!("invalid IPv4 address {}/{}: {:?}", address, prefix, err).into() + }) +} + +fn ipv6_inet(address: Ipv6Addr, prefix: i32) -> Result { + let prefix = u8::try_from(prefix) + .map_err(|_| anyhow::anyhow!("invalid IPv6 prefix length {}", prefix))?; + Ipv6Inet::new(address, prefix).map_err(|err| { + anyhow::anyhow!("invalid IPv6 address {}/{}: {:?}", address, prefix, err).into() + }) +} + +#[cfg(test)] +mod tests { + use std::str::FromStr as _; + + use crate::common::{ifcfg::IfConfiguerTrait, netns::NetNS}; + use tokio::sync::Notify; + + use super::*; + + struct FailingRemoveIpIfConfiger; + + #[async_trait::async_trait] + impl IfConfiguerTrait for FailingRemoveIpIfConfiger { + async fn remove_ip(&self, _name: &str, _ip: Option) -> Result<(), Error> { + Err(anyhow::anyhow!("forced remove_ip failure").into()) + } + } + + fn member_id(n: u128) -> SharedVirtualNicMemberId { + uuid::Uuid::from_u128(n) + } + + fn claims_with_ipv4_route(route: SharedIpv4Route, mtu: Option) -> SharedIfConfigClaims { + SharedIfConfigClaims { + ipv4_routes: BTreeSet::from([route]), + mtu, + ..Default::default() + } + } + + fn claims_with_ipv4_address_and_route( + address: Ipv4Inet, + route: SharedIpv4Route, + ) -> SharedIfConfigClaims { + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([address]), + ipv4_routes: BTreeSet::from([route]), + ..Default::default() + } + } + + fn virtual_nic_config() -> VirtualNicConfig { + VirtualNicConfig::new(String::new(), 1500, NetNS::new(None)) + } + + fn virtual_nic_config_with_mtu(mtu: u32) -> VirtualNicConfig { + VirtualNicConfig::new(String::new(), mtu, NetNS::new(None)) + } + + fn virtual_nic_config_in_netns(net_ns: &str) -> VirtualNicConfig { + VirtualNicConfig::new(String::new(), 1500, NetNS::new(Some(net_ns.to_owned()))) + } + + #[test] + fn duplicate_routes_keep_owner_sets_and_single_os_delta() { + let route = SharedIpv4Route::new(Ipv4Addr::new(100, 100, 100, 101), 32, None); + let first = member_id(1); + let second = member_id(2); + let mut ifcfg = SharedIfConfig::default(); + + let first_delta = + ifcfg.apply_member_claims(first, claims_with_ipv4_route(route.clone(), Some(1400))); + let second_delta = + ifcfg.apply_member_claims(second, claims_with_ipv4_route(route.clone(), Some(1300))); + + assert_eq!( + first_delta.ipv4_routes.added, + BTreeSet::from([route.clone()]) + ); + assert!(second_delta.ipv4_routes.added.is_empty()); + assert_eq!( + ifcfg.owners_of_ipv4_route(&route), + BTreeSet::from([first, second]) + ); + assert_eq!(ifcfg.effective_mtu(), Some(1300)); + } + + #[test] + fn removing_one_owner_keeps_shared_route_until_last_owner_leaves() { + let route = SharedIpv4Route::new(Ipv4Addr::new(100, 100, 100, 101), 32, None); + let first = member_id(1); + let second = member_id(2); + let mut ifcfg = SharedIfConfig::default(); + ifcfg.apply_member_claims(first, claims_with_ipv4_route(route.clone(), None)); + ifcfg.apply_member_claims(second, claims_with_ipv4_route(route.clone(), None)); + + let first_delta = ifcfg.remove_member(first).unwrap(); + let second_delta = ifcfg.remove_member(second).unwrap(); + + assert!(first_delta.ipv4_routes.removed.is_empty()); + assert_eq!( + second_delta.ipv4_routes.removed, + BTreeSet::from([route.clone()]) + ); + assert!(ifcfg.owners_of_ipv4_route(&route).is_empty()); + } + + #[tokio::test] + async fn overlapping_member_claims_leave_existing_configuration_intact() { + let first = member_id(1); + let second = member_id(2); + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + shared_nic.ifcfg_mut().apply_member_claims( + first, + SharedIfConfigClaims { + ipv4_routes: BTreeSet::from([SharedIpv4Route::new( + Ipv4Addr::new(10, 144, 0, 0), + 16, + None, + )]), + ..Default::default() + }, + ); + let original = shared_nic.ifcfg().effective(); + + let err = shared_nic + .apply_member_claims( + second, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from(["10.144.1.2/24".parse().unwrap()]), + ..Default::default() + }, + ) + .await + .unwrap_err(); + + assert!(err.to_string().contains("overlaps member")); + assert_eq!(shared_nic.ifcfg().effective(), original); + } + + #[test] + fn shared_magic_dns_route_is_not_a_member_conflict() { + let first = member_id(1); + let second = member_id(2); + let route = SharedIpv4Route::new(Ipv4Addr::new(100, 100, 100, 101), 32, None); + let claims = claims_with_ipv4_route(route, None); + let mut ifcfg = SharedIfConfig::default(); + ifcfg.apply_member_claims(first, claims.clone()); + + assert!(ifcfg.ensure_disjoint(second, &claims).is_ok()); + } + + #[test] + fn member_claim_update_tracks_ip_ownership() { + let first_ip = Ipv4Inet::from_str("10.30.0.2/24").unwrap(); + let second_ip = Ipv4Inet::from_str("10.30.0.3/24").unwrap(); + let member = member_id(1); + let mut ifcfg = SharedIfConfig::default(); + + let first_delta = ifcfg.apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([first_ip]), + ..Default::default() + }, + ); + let second_delta = ifcfg.apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([second_ip]), + ..Default::default() + }, + ); + + assert_eq!(first_delta.ipv4_addresses.added, BTreeSet::from([first_ip])); + assert_eq!( + second_delta.ipv4_addresses.removed, + BTreeSet::from([first_ip]) + ); + assert_eq!( + second_delta.ipv4_addresses.added, + BTreeSet::from([second_ip]) + ); + assert_eq!( + ifcfg.owners_of_ipv4_address(&second_ip), + BTreeSet::from([member]) + ); + } + + #[test] + fn ipv4_route_source_hint_prefers_address_inside_route() { + let route = SharedIpv4Route::new(Ipv4Addr::new(10, 90, 1, 0), 24, None); + let member = member_id(1); + let mut ifcfg = SharedIfConfig::default(); + + ifcfg.apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([ + Ipv4Inet::from_str("10.1.1.1/24").unwrap(), + Ipv4Inet::from_str("10.90.1.1/24").unwrap(), + ]), + ipv4_routes: BTreeSet::from([route.clone()]), + ..Default::default() + }, + ); + + assert_eq!( + ifcfg.ipv4_route_source_hint(&route), + Some(Ipv4Addr::new(10, 90, 1, 1)) + ); + } + + #[test] + fn adding_better_ipv4_route_owner_marks_source_change() { + let route = SharedIpv4Route::new(Ipv4Addr::new(10, 90, 2, 0), 24, None); + let first = member_id(1); + let second = member_id(2); + let mut ifcfg = SharedIfConfig::default(); + ifcfg.apply_member_claims( + first, + claims_with_ipv4_address_and_route( + Ipv4Inet::from_str("10.1.2.1/24").unwrap(), + route.clone(), + ), + ); + let old = ifcfg.clone(); + + let delta = ifcfg.apply_member_claims( + second, + claims_with_ipv4_address_and_route( + Ipv4Inet::from_str("10.90.2.1/24").unwrap(), + route.clone(), + ), + ); + + assert!(delta.ipv4_routes.added.is_empty()); + assert_eq!( + old.changed_ipv4_route_sources(&ifcfg), + BTreeSet::from([route.clone()]) + ); + assert_eq!( + old.ipv4_route_source_hint(&route), + Some(Ipv4Addr::new(10, 1, 2, 1)) + ); + assert_eq!( + ifcfg.ipv4_route_source_hint(&route), + Some(Ipv4Addr::new(10, 90, 2, 1)) + ); + } + + #[test] + fn removing_ipv4_route_owner_marks_source_change_when_route_remains() { + let route = SharedIpv4Route::new(Ipv4Addr::new(10, 90, 3, 0), 24, None); + let first = member_id(1); + let second = member_id(2); + let mut ifcfg = SharedIfConfig::default(); + ifcfg.apply_member_claims( + first, + claims_with_ipv4_address_and_route( + Ipv4Inet::from_str("10.1.3.1/24").unwrap(), + route.clone(), + ), + ); + ifcfg.apply_member_claims( + second, + claims_with_ipv4_address_and_route( + Ipv4Inet::from_str("10.90.3.1/24").unwrap(), + route.clone(), + ), + ); + let old = ifcfg.clone(); + + let delta = ifcfg.remove_member(second).unwrap(); + + assert!(delta.ipv4_routes.removed.is_empty()); + assert_eq!( + old.changed_ipv4_route_sources(&ifcfg), + BTreeSet::from([route.clone()]) + ); + assert_eq!( + old.ipv4_route_source_hint(&route), + Some(Ipv4Addr::new(10, 90, 3, 1)) + ); + assert_eq!( + ifcfg.ipv4_route_source_hint(&route), + Some(Ipv4Addr::new(10, 1, 3, 1)) + ); + } + + #[test] + fn removing_ipv4_route_keeps_old_source_hint_available() { + let kept_route = SharedIpv4Route::new(Ipv4Addr::new(10, 90, 4, 0), 24, Some(10)); + let removed_route = SharedIpv4Route::new(Ipv4Addr::new(10, 90, 4, 0), 24, Some(20)); + let member = member_id(1); + let address = Ipv4Inet::from_str("10.90.4.1/24").unwrap(); + let mut ifcfg = SharedIfConfig::default(); + ifcfg.apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([address]), + ipv4_routes: BTreeSet::from([kept_route.clone(), removed_route.clone()]), + ..Default::default() + }, + ); + let old = ifcfg.clone(); + + let delta = ifcfg.apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([address]), + ipv4_routes: BTreeSet::from([kept_route.clone()]), + ..Default::default() + }, + ); + + assert_eq!( + delta.ipv4_routes.removed, + BTreeSet::from([removed_route.clone()]) + ); + assert_eq!( + old.ipv4_route_source_hint(&removed_route), + Some(Ipv4Addr::new(10, 90, 4, 1)) + ); + } + + #[test] + fn shared_virtual_nic_wraps_virtual_nic_and_tracks_ifcfg() { + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + let member = member_id(1); + let route = SharedIpv4Route::new(Ipv4Addr::new(10, 40, 0, 0), 24, None); + + shared_nic + .ifcfg_mut() + .apply_member_claims(member, claims_with_ipv4_route(route.clone(), None)); + + assert_eq!( + shared_nic.ifcfg().owners_of_ipv4_route(&route), + BTreeSet::from([member]) + ); + drop(shared_nic.nic()); + } + + #[tokio::test] + async fn stale_member_registration_cleanup_keeps_current_claims() { + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + let member = member_id(1); + let old_registration = uuid::Uuid::from_u128(10); + let current_registration = uuid::Uuid::from_u128(11); + let ip = Ipv4Inet::from_str("10.50.0.2/24").unwrap(); + + shared_nic + .member_registrations + .insert(member, current_registration); + shared_nic.ifcfg_mut().apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([ip]), + ..Default::default() + }, + ); + + shared_nic + .remove_member_registration_claims(member, old_registration) + .await + .unwrap(); + + assert_eq!( + shared_nic.ifcfg().owners_of_ipv4_address(&ip), + BTreeSet::from([member]) + ); + assert_eq!( + shared_nic.member_registrations.get(&member), + Some(¤t_registration) + ); + } + + #[tokio::test] + async fn failed_member_registration_cleanup_invalidates_shared_nic() { + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + let member = member_id(1); + let registration = uuid::Uuid::from_u128(10); + let ip = Ipv4Inet::from_str("10.60.0.2/24").unwrap(); + + shared_nic.member_registrations.insert(member, registration); + shared_nic.ifcfg_mut().apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([ip]), + ..Default::default() + }, + ); + let nic = shared_nic.nic(); + let mut nic = nic.lock().await; + nic.set_ifname_for_test("et0".to_string()); + nic.set_ifcfg_for_test(Box::new(FailingRemoveIpIfConfiger)); + drop(nic); + + let result = shared_nic + .remove_member_registration_claims(member, registration) + .await; + + assert!(result.is_err()); + assert!(!shared_nic.is_valid()); + assert!(!shared_nic.member_registrations.contains_key(&member)); + assert_eq!( + shared_nic.ifcfg().owners_of_ipv4_address(&ip), + BTreeSet::from([member]) + ); + } + + #[tokio::test] + async fn failed_registration_replacement_keeps_old_registration_and_invalidates() { + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + let member = member_id(1); + let old_registration = uuid::Uuid::from_u128(10); + let next_registration = uuid::Uuid::from_u128(11); + let ip = Ipv4Inet::from_str("10.70.0.2/24").unwrap(); + + shared_nic + .member_registrations + .insert(member, old_registration); + shared_nic.ifcfg_mut().apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([ip]), + ..Default::default() + }, + ); + let nic = shared_nic.nic(); + let mut nic = nic.lock().await; + nic.set_ifname_for_test("et0".to_string()); + nic.set_ifcfg_for_test(Box::new(FailingRemoveIpIfConfiger)); + drop(nic); + + let result = shared_nic + .attach_member_registration(member, next_registration) + .await; + + assert!(result.is_err()); + assert!(!shared_nic.is_valid()); + assert_eq!( + shared_nic.member_registrations.get(&member), + Some(&old_registration) + ); + } + + #[test] + fn registry_reuses_shared_virtual_nic_for_same_dev_name_and_netns() { + let mut registry = SharedVirtualNicRegistry::new(); + + let first = registry.get_or_create("et0".to_string(), virtual_nic_config()); + let second = registry.get_or_create("et0".to_string(), virtual_nic_config()); + + assert!(Arc::ptr_eq(&first, &second)); + } + + #[test] + fn registry_keeps_same_dev_name_in_different_netns_separate() { + let mut registry = SharedVirtualNicRegistry::new(); + + let first = registry.get_or_create("et0".to_string(), virtual_nic_config_in_netns("net-a")); + let second = + registry.get_or_create("et0".to_string(), virtual_nic_config_in_netns("net-b")); + + assert!(!Arc::ptr_eq(&first, &second)); + } + + #[test] + fn registry_keeps_different_dev_names_separate() { + let mut registry = SharedVirtualNicRegistry::new(); + + let first = registry.get_or_create("et0".to_string(), virtual_nic_config()); + let second = registry.get_or_create("et1".to_string(), virtual_nic_config()); + + assert!(!Arc::ptr_eq(&first, &second)); + } + + #[test] + fn registry_replaces_invalid_shared_virtual_nic() { + let mut registry = SharedVirtualNicRegistry::new(); + + let first = registry.get_or_create("et0".to_string(), virtual_nic_config()); + first.try_lock().unwrap().mark_invalid(); + let second = registry.get_or_create("et0".to_string(), virtual_nic_config()); + + assert!(!Arc::ptr_eq(&first, &second)); + assert!( + registry + .get("et0", &virtual_nic_config()) + .is_some_and(|nic| Arc::ptr_eq(&nic, &second)) + ); + } + + #[test] + fn registry_create_member_uses_registered_shared_virtual_nic() { + let mut registry = SharedVirtualNicRegistry::new(); + let member_id = member_id(1); + + let member = registry.create_member( + "et0".to_string(), + virtual_nic_config(), + member_id, + Arc::new(Notify::new()), + ); + let shared_nic = registry.get("et0", &virtual_nic_config()).unwrap(); + + assert_eq!(member.member_id(), member_id); + assert!(Arc::ptr_eq(&member.shared_nic(), &shared_nic)); + } + + #[test] + fn registry_create_member_keeps_member_configured_mtu() { + let mut registry = SharedVirtualNicRegistry::new(); + + let first = registry.create_member( + "et0".to_string(), + virtual_nic_config_with_mtu(1400), + member_id(1), + Arc::new(Notify::new()), + ); + let second = registry.create_member( + "et0".to_string(), + virtual_nic_config_with_mtu(1300), + member_id(2), + Arc::new(Notify::new()), + ); + + assert_eq!(first.configured_mtu_for_test(), 1400); + assert_eq!(second.configured_mtu_for_test(), 1300); + assert!(Arc::ptr_eq(&first.shared_nic(), &second.shared_nic())); + } +} diff --git a/easytier/src/instance/shared_virtual_nic/dispatcher.rs b/easytier/src/instance/shared_virtual_nic/dispatcher.rs new file mode 100644 index 00000000..a766f745 --- /dev/null +++ b/easytier/src/instance/shared_virtual_nic/dispatcher.rs @@ -0,0 +1,2657 @@ +use std::{ + collections::{BTreeMap, BTreeSet, HashMap, VecDeque}, + pin::Pin, + sync::{ + Arc, Mutex as StdMutex, + atomic::{AtomicBool, Ordering}, + }, + time::Duration, +}; + +use cidr::{Ipv4Inet, Ipv6Inet}; +use futures::{SinkExt, StreamExt}; +use pnet_packet::{ + MutablePacket as _, Packet as _, + icmp::{self, MutableIcmpPacket}, + ip::IpNextHeaderProtocols, + ipv4::{self, MutableIpv4Packet}, + ipv6::MutableIpv6Packet, + tcp::{self, MutableTcpPacket}, + udp::{self, MutableUdpPacket}, +}; +#[cfg(mobile)] +use std::sync::OnceLock; +#[cfg(mobile)] +use tokio::runtime::{Builder, Runtime}; +#[cfg(mobile)] +use tokio::sync::Mutex; +use tokio::sync::{Notify, mpsc, oneshot}; +use tokio_util::task::AbortOnDropHandle; + +use crate::common::error::Error; +#[cfg(mobile)] +use crate::instance::virtual_nic::VirtualNic; +use easytier_core::{ + packet::ZCPacket, + tunnel::{Tunnel, ZCPacketSink, ZCPacketStream}, +}; + +use super::{ + SharedIfConfigClaims, SharedIpv4Route, SharedIpv6Route, SharedVirtualNicMemberId, + SharedVirtualNicMemberRegistrationId, +}; + +const MEMBER_TUNNEL_BUFFER_SIZE: usize = 1024; +const FLOW_OWNER_LIMIT: usize = 4096; +const IPV4_HEADER_MIN_LEN: usize = 20; +const IPV6_HEADER_LEN: usize = 40; +const TCP_HEADER_MIN_LEN: usize = 20; +const UDP_HEADER_LEN: usize = 8; +const ICMP_ECHO_HEADER_LEN: usize = 8; +const ICMP_PROTOCOL: u8 = 1; +const ICMPV6_PROTOCOL: u8 = 58; +const TCP_PROTOCOL: u8 = 6; +const UDP_PROTOCOL: u8 = 17; +const DISPATCHER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(1); + +#[cfg(mobile)] +fn mobile_dispatcher_runtime() -> &'static Runtime { + static RUNTIME: OnceLock = OnceLock::new(); + RUNTIME.get_or_init(|| { + Builder::new_multi_thread() + .worker_threads(1) + .thread_name("easytier-shared-tun") + .enable_all() + .build() + .expect("failed to build shared virtual nic mobile dispatcher runtime") + }) +} + +struct SharedVirtualNicMemberPacket { + member_id: SharedVirtualNicMemberId, + packet: ZCPacket, +} + +#[cfg(mobile)] +struct SharedVirtualNicMobileTunUpdate { + tun_fd: std::os::fd::RawFd, + ack: oneshot::Sender>, +} + +enum SharedVirtualNicControl { + Register { + member_id: SharedVirtualNicMemberId, + entry: SharedVirtualNicMemberTunnelEntry, + }, + Unregister { + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + }, + UpdateSources { + member_id: SharedVirtualNicMemberId, + sources: SharedVirtualNicMemberSources, + ack: oneshot::Sender<()>, + }, + Shutdown { + invalidate: bool, + ack: oneshot::Sender<()>, + }, +} + +#[derive(Clone, Default)] +pub(super) struct SharedVirtualNicMemberTunnelTable { + state: Arc>, +} + +#[derive(Default)] +struct SharedVirtualNicMemberTunnelTableState { + to_tun_sender: Option>, + control_sender: Option>, +} + +struct SharedVirtualNicMemberTunnelEntry { + registration_id: SharedVirtualNicMemberRegistrationId, + sender: mpsc::Sender, + close_notifier: Arc, + _tasks: Vec>, +} + +impl SharedVirtualNicMemberTunnelTable { + fn attach_dispatcher( + &self, + to_tun_sender: mpsc::Sender, + control_sender: mpsc::UnboundedSender, + ) { + let mut state = self.state.lock().unwrap(); + state.to_tun_sender = Some(to_tun_sender); + state.control_sender = Some(control_sender); + } + + pub(super) fn detach_dispatcher(&self) { + let mut state = self.state.lock().unwrap(); + state.to_tun_sender.take(); + state.control_sender.take(); + } + + pub(super) fn register( + &self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + tunnel: Box, + close_notifier: Arc, + ) -> Result<(), Error> { + let channels = self + .dispatcher_channels() + .ok_or_else(|| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; + let (to_tun_sender, control_sender) = channels; + let (mut member_stream, mut member_sink) = tunnel.split(); + let (to_member_sender, mut to_member_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (reader_start_sender, reader_start_receiver) = oneshot::channel(); + + let reader_control_sender = control_sender.clone(); + let reader_close_notifier = close_notifier.clone(); + let reader_task = AbortOnDropHandle::new(tokio::spawn(async move { + if reader_start_receiver.await.is_err() { + return; + } + + while let Some(packet) = member_stream.next().await { + let packet = match packet { + Ok(packet) => packet, + Err(err) => { + tracing::error!(?member_id, ?err, "shared member tunnel read failed"); + break; + } + }; + + if to_tun_sender + .send(SharedVirtualNicMemberPacket { member_id, packet }) + .await + .is_err() + { + break; + } + } + + notify_member_tunnel_closed( + &reader_control_sender, + &reader_close_notifier, + member_id, + registration_id, + ); + })); + + let writer_control_sender = control_sender.clone(); + let writer_close_notifier = close_notifier.clone(); + let writer_task = AbortOnDropHandle::new(tokio::spawn(async move { + while let Some(packet) = to_member_receiver.recv().await { + if let Err(err) = member_sink.send(packet).await { + tracing::error!(?member_id, ?err, "shared member tunnel write failed"); + notify_member_tunnel_closed( + &writer_control_sender, + &writer_close_notifier, + member_id, + registration_id, + ); + break; + } + } + })); + + let entry = SharedVirtualNicMemberTunnelEntry { + registration_id, + sender: to_member_sender, + close_notifier, + _tasks: vec![reader_task, writer_task], + }; + control_sender + .send(SharedVirtualNicControl::Register { member_id, entry }) + .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; + let _ = reader_start_sender.send(()); + + Ok(()) + } + + pub(super) fn unregister( + &self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) { + let Some(control_sender) = self.control_sender() else { + return; + }; + let _ = control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); + } + + fn dispatcher_channels( + &self, + ) -> Option<( + mpsc::Sender, + mpsc::UnboundedSender, + )> { + let state = self.state.lock().unwrap(); + Some((state.to_tun_sender.clone()?, state.control_sender.clone()?)) + } + + fn control_sender(&self) -> Option> { + self.state.lock().unwrap().control_sender.clone() + } +} + +fn notify_member_tunnel_closed( + control_sender: &mpsc::UnboundedSender, + close_notifier: &Notify, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, +) { + let _ = control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); + close_notifier.notify_one(); +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] +enum SharedVirtualNicFlowAddr { + V4(u32), + V6([u8; 16]), +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +struct SharedVirtualNicTransportPorts { + src: u16, + dst: u16, +} + +impl SharedVirtualNicTransportPorts { + fn reversed(self) -> Self { + Self { + src: self.dst, + dst: self.src, + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +struct SharedVirtualNicFlowKey { + src: SharedVirtualNicFlowAddr, + dst: SharedVirtualNicFlowAddr, + protocol: u8, + ports: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +struct SharedVirtualNicMemberSources { + ipv4_addresses: BTreeSet, + ipv6_addresses: BTreeSet, + ipv4_routes: BTreeSet, + ipv6_routes: BTreeSet, +} + +impl SharedVirtualNicMemberSources { + fn from_claims(claims: &SharedIfConfigClaims) -> Self { + let ipv4_addresses = claims + .ipv4_addresses + .iter() + .filter(|addr| !addr.address().is_unspecified()) + .copied() + .collect::>(); + let ipv6_addresses = claims + .ipv6_addresses + .iter() + .filter(|addr| !addr.address().is_unspecified()) + .copied() + .collect::>(); + let ipv4_routes = claims + .ipv4_routes + .iter() + .filter_map(ipv4_route_to_inet) + .collect::>(); + let ipv6_routes = claims + .ipv6_routes + .iter() + .filter_map(ipv6_route_to_inet) + .collect::>(); + + Self { + ipv4_addresses, + ipv6_addresses, + ipv4_routes, + ipv6_routes, + } + } + + fn is_empty(&self) -> bool { + self.ipv4_addresses.is_empty() + && self.ipv6_addresses.is_empty() + && self.ipv4_routes.is_empty() + && self.ipv6_routes.is_empty() + } + + fn owns_source(&self, source: SharedVirtualNicFlowAddr) -> bool { + match source { + SharedVirtualNicFlowAddr::V4(source) => self + .ipv4_addresses + .iter() + .any(|address| u32::from(address.address()) == source), + SharedVirtualNicFlowAddr::V6(source) => self + .ipv6_addresses + .iter() + .any(|address| address.address().octets() == source), + } + } +} + +fn ipv4_route_to_inet(route: &SharedIpv4Route) -> Option { + Ipv4Inet::new(route.address, route.prefix).ok() +} + +fn ipv6_route_to_inet(route: &SharedIpv6Route) -> Option { + Ipv6Inet::new(route.address, route.prefix).ok() +} + +impl From for SharedVirtualNicFlowAddr { + fn from(addr: std::net::Ipv4Addr) -> Self { + Self::V4(u32::from_be_bytes(addr.octets())) + } +} + +impl From for SharedVirtualNicFlowAddr { + fn from(addr: std::net::Ipv6Addr) -> Self { + Self::V6(addr.octets()) + } +} + +impl SharedVirtualNicFlowAddr { + fn as_ipv4(self) -> Option { + match self { + Self::V4(addr) => Some(std::net::Ipv4Addr::from(addr)), + Self::V6(_) => None, + } + } + + fn as_ipv6(self) -> Option { + match self { + Self::V4(_) => None, + Self::V6(addr) => Some(std::net::Ipv6Addr::from(addr)), + } + } + fn packet_addrs(packet: &ZCPacket) -> Option<(Self, Self)> { + let payload = packet.payload(); + match payload.first()? >> 4 { + 4 => SharedVirtualNicFlowKey::from_ipv4_payload(payload).map(|key| (key.src, key.dst)), + 6 if payload.len() >= IPV6_HEADER_LEN => Some(( + Self::V6(read_ipv6_addr(payload, 8)), + Self::V6(read_ipv6_addr(payload, 24)), + )), + _ => None, + } + } +} + +impl SharedVirtualNicFlowKey { + fn from_packet(packet: &ZCPacket) -> Option { + let payload = packet.payload(); + let version = payload.first()? >> 4; + match version { + 4 => Self::from_ipv4_payload(payload), + 6 => Self::from_ipv6_payload(payload), + _ => None, + } + } + + fn from_ipv4_payload(payload: &[u8]) -> Option { + if payload.len() < IPV4_HEADER_MIN_LEN { + return None; + } + + let header_len = usize::from(payload[0] & 0x0f) * 4; + if header_len < IPV4_HEADER_MIN_LEN || payload.len() < header_len { + return None; + } + + let protocol = payload[9]; + let fragment_offset = u16::from_be_bytes([payload[6], payload[7]]) & 0x1fff; + let src = u32::from_be_bytes([payload[12], payload[13], payload[14], payload[15]]); + let dst = u32::from_be_bytes([payload[16], payload[17], payload[18], payload[19]]); + Some(Self { + src: SharedVirtualNicFlowAddr::V4(src), + dst: SharedVirtualNicFlowAddr::V4(dst), + protocol, + ports: if fragment_offset == 0 { + transport_ports(protocol, &payload[header_len..]) + } else { + None + }, + }) + } + + fn from_ipv6_payload(payload: &[u8]) -> Option { + if payload.len() < IPV6_HEADER_LEN || payload[0] >> 4 != 6 { + return None; + } + + let protocol = payload[6]; + Some(Self { + src: SharedVirtualNicFlowAddr::V6(read_ipv6_addr(payload, 8)), + dst: SharedVirtualNicFlowAddr::V6(read_ipv6_addr(payload, 24)), + protocol, + ports: transport_ports(protocol, &payload[IPV6_HEADER_LEN..]), + }) + } + + fn reversed(&self) -> Self { + Self { + src: self.dst, + dst: self.src, + protocol: self.protocol, + ports: self.ports.map(|ports| { + if matches!(self.protocol, ICMP_PROTOCOL | ICMPV6_PROTOCOL) { + ports + } else { + ports.reversed() + } + }), + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct SharedVirtualNicNatEntry { + original_src: SharedVirtualNicFlowAddr, + translated_src: SharedVirtualNicFlowAddr, +} + +#[derive(Default)] +struct SharedVirtualNicNatTable { + entries: HashMap, + insert_order: VecDeque, +} + +impl SharedVirtualNicNatTable { + fn remember( + &mut self, + translated_packet: &ZCPacket, + original_src: SharedVirtualNicFlowAddr, + translated_src: SharedVirtualNicFlowAddr, + ) { + let Some(key) = + SharedVirtualNicFlowKey::from_packet(translated_packet).map(|key| key.reversed()) + else { + return; + }; + + if !self.entries.contains_key(&key) { + self.evict_before_insert(); + self.insert_order.push_back(key); + } + self.entries.insert( + key, + SharedVirtualNicNatEntry { + original_src, + translated_src, + }, + ); + } + + fn translate_reply(&mut self, packet: &mut ZCPacket) -> bool { + let Some(key) = SharedVirtualNicFlowKey::from_packet(packet) else { + return false; + }; + let Some(entry) = self.entries.get(&key).copied() else { + return false; + }; + + rewrite_packet_destination(packet, entry.translated_src, entry.original_src) + } + + fn clear(&mut self) { + self.entries.clear(); + self.insert_order.clear(); + } + + fn evict_before_insert(&mut self) { + while self.entries.len() >= FLOW_OWNER_LIMIT { + let Some(key) = self.insert_order.pop_front() else { + self.entries.clear(); + return; + }; + self.entries.remove(&key); + } + } +} + +pub(super) struct SharedVirtualNicDispatcher { + _task: AbortOnDropHandle<()>, + control_sender: mpsc::UnboundedSender, + #[cfg(mobile)] + mobile_tun_update_sender: Option>, +} + +impl SharedVirtualNicDispatcher { + pub(super) fn start( + tunnel: Box, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + ) -> Self { + let (tun_stream, tun_sink) = tunnel.split(); + let (to_tun_sender, to_tun_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (control_sender, control_receiver) = mpsc::unbounded_channel(); + member_tunnel_table.attach_dispatcher(to_tun_sender, control_sender.clone()); + + let task = SharedVirtualNicDispatcherTask { + tun_stream, + tun_sink, + to_tun_receiver, + control_receiver, + member_tunnel_table, + valid, + state: SharedVirtualNicDispatcherState::default(), + }; + + Self { + _task: AbortOnDropHandle::new(tokio::spawn(task.run())), + control_sender, + #[cfg(mobile)] + mobile_tun_update_sender: None, + } + } + + #[cfg(mobile)] + pub(super) async fn start_for_mobile( + nic: Arc>, + tun_fd: std::os::fd::RawFd, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + ) -> Result { + let (to_tun_sender, to_tun_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (control_sender, control_receiver) = mpsc::unbounded_channel(); + let (mobile_tun_update_sender, mobile_tun_update_receiver) = mpsc::unbounded_channel(); + let (first_open_sender, first_open_receiver) = oneshot::channel(); + let task_member_tunnel_table = member_tunnel_table.clone(); + member_tunnel_table.attach_dispatcher(to_tun_sender, control_sender.clone()); + + let task = mobile_dispatcher_runtime().spawn(async move { + let task = SharedVirtualNicMobileDispatcherTask { + nic, + tun_stream: None, + tun_sink: None, + tun_fd, + tun_update_receiver: mobile_tun_update_receiver, + to_tun_receiver, + control_receiver, + member_tunnel_table: task_member_tunnel_table, + valid, + state: SharedVirtualNicDispatcherState::default(), + }; + + task.run(first_open_sender).await; + }); + + let dispatcher = Self { + _task: AbortOnDropHandle::new(task), + control_sender, + mobile_tun_update_sender: Some(mobile_tun_update_sender), + }; + await_mobile_open(first_open_receiver) + .await + .map(|()| dispatcher) + } + + #[cfg(mobile)] + pub(super) async fn update_mobile_tun_fd( + &self, + tun_fd: std::os::fd::RawFd, + ) -> Result<(), Error> { + let sender = self.mobile_tun_update_sender.as_ref().ok_or_else(|| { + Error::from(anyhow::anyhow!( + "shared virtual nic mobile dispatcher is not running" + )) + })?; + let (ack, receiver) = oneshot::channel(); + sender + .send(SharedVirtualNicMobileTunUpdate { tun_fd, ack }) + .map_err(|_| { + Error::from(anyhow::anyhow!( + "shared virtual nic mobile dispatcher is not running" + )) + })?; + await_mobile_open(receiver).await + } + + pub(super) async fn update_sources( + &self, + member_id: SharedVirtualNicMemberId, + claims: &SharedIfConfigClaims, + ) -> Result<(), Error> { + self.send_source_update( + member_id, + SharedVirtualNicMemberSources::from_claims(claims), + ) + .await + } + + pub(super) async fn remove_sources( + &self, + member_id: SharedVirtualNicMemberId, + ) -> Result<(), Error> { + self.send_source_update(member_id, SharedVirtualNicMemberSources::default()) + .await + } + + async fn send_source_update( + &self, + member_id: SharedVirtualNicMemberId, + sources: SharedVirtualNicMemberSources, + ) -> Result<(), Error> { + let (ack, rx) = oneshot::channel(); + self.control_sender + .send(SharedVirtualNicControl::UpdateSources { + member_id, + sources, + ack, + }) + .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; + rx.await + .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running").into()) + } + + pub(super) async fn shutdown_without_invalidation(self) { + let (ack, rx) = oneshot::channel(); + if self + .control_sender + .send(SharedVirtualNicControl::Shutdown { + invalidate: false, + ack, + }) + .is_err() + { + return; + } + + if tokio::time::timeout(DISPATCHER_SHUTDOWN_TIMEOUT, rx) + .await + .is_err() + { + tracing::warn!("timed out shutting down shared virtual nic dispatcher"); + } + } +} + +enum DispatcherControlResult { + Continue, + Stop { + invalidate: bool, + ack: Option>, + }, +} + +fn handle_dispatcher_control( + state: &mut SharedVirtualNicDispatcherState, + control: Option, +) -> DispatcherControlResult { + let Some(control) = control else { + return DispatcherControlResult::Stop { + invalidate: true, + ack: None, + }; + }; + + match control { + SharedVirtualNicControl::Shutdown { invalidate, ack } => DispatcherControlResult::Stop { + invalidate, + ack: Some(ack), + }, + other => { + state.handle_control(other); + DispatcherControlResult::Continue + } + } +} + +fn acknowledge_dispatcher_shutdown(ack: Option>) { + if let Some(ack) = ack { + let _ = ack.send(()); + } +} + +struct SharedVirtualNicDispatcherTask { + tun_stream: Pin>, + tun_sink: Pin>, + to_tun_receiver: mpsc::Receiver, + control_receiver: mpsc::UnboundedReceiver, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + state: SharedVirtualNicDispatcherState, +} + +impl SharedVirtualNicDispatcherTask { + async fn run(mut self) { + loop { + tokio::select! { + control = self.control_receiver.recv() => { + if let DispatcherControlResult::Stop { invalidate, ack } = + handle_dispatcher_control(&mut self.state, control) + { + self.cleanup(invalidate); + acknowledge_dispatcher_shutdown(ack); + return; + } + } + member_packet = self.to_tun_receiver.recv() => { + let Some(member_packet) = member_packet else { + break; + }; + if !self.forward_member_packet_to_tun(member_packet).await { + break; + } + } + packet = self.tun_stream.next() => { + let Some(packet) = packet else { + break; + }; + let packet = match packet { + Ok(packet) => packet, + Err(err) => { + tracing::error!(?err, "shared virtual nic read from tun failed"); + break; + } + }; + self.state.forward_tun_packet_to_member(packet).await; + } + } + } + + self.cleanup(true); + } + + fn cleanup(&mut self, invalidate: bool) { + if invalidate { + self.valid.store(false, Ordering::Release); + } + self.member_tunnel_table.detach_dispatcher(); + self.state.close_all(); + } + + async fn forward_member_packet_to_tun( + &mut self, + member_packet: SharedVirtualNicMemberPacket, + ) -> bool { + let packet = self + .state + .prepare_member_packet_to_tun(member_packet.member_id, member_packet.packet); + if let Err(err) = self.tun_sink.send(packet).await { + tracing::error!(?err, "shared virtual nic write to tun failed"); + return false; + } + true + } +} + +#[cfg(mobile)] +struct SharedVirtualNicMobileDispatcherTask { + nic: Arc>, + tun_stream: Option>>, + tun_sink: Option>>, + tun_fd: std::os::fd::RawFd, + tun_update_receiver: mpsc::UnboundedReceiver, + to_tun_receiver: mpsc::Receiver, + control_receiver: mpsc::UnboundedReceiver, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + state: SharedVirtualNicDispatcherState, +} + +#[cfg(mobile)] +impl SharedVirtualNicMobileDispatcherTask { + async fn run(mut self, first_open_sender: oneshot::Sender>) { + match self.open_tun().await { + Ok(()) => { + if first_open_sender.send(Ok(())).is_err() { + self.cleanup(false); + return; + } + } + Err(err) => { + self.cleanup(false); + let _ = first_open_sender.send(Err(err)); + return; + } + } + + loop { + tokio::select! { + control = self.control_receiver.recv() => { + if !self.handle_control(control) { + return; + } + } + member_packet = self.to_tun_receiver.recv() => { + let Some(member_packet) = member_packet else { + self.cleanup(true); + return; + }; + self.forward_member_packet_to_tun(member_packet).await; + } + packet = async { + match self.tun_stream.as_mut() { + Some(stream) => stream.next().await, + None => std::future::pending().await, + } + } => { + let Some(packet) = packet else { + tracing::error!("shared virtual nic mobile tun stream closed"); + self.drop_tun(); + continue; + }; + match packet { + Ok(packet) => self.state.forward_tun_packet_to_member(packet).await, + Err(err) => { + tracing::error!(?err, "shared virtual nic read from mobile tun failed"); + self.drop_tun(); + } + } + } + update = self.tun_update_receiver.recv() => { + let Some(update) = update else { + self.cleanup(true); + return; + }; + self.apply_tun_update(update).await; + } + } + } + } + + async fn open_tun(&mut self) -> Result<(), Error> { + let tun_fd = self.tun_fd; + let tunnel = self.nic.lock().await.create_dev_for_mobile(tun_fd).await?; + let (tun_stream, tun_sink) = tunnel.split(); + self.tun_stream = Some(tun_stream); + self.tun_sink = Some(tun_sink); + tracing::info!(fd = tun_fd, "opened shared virtual nic mobile tun"); + Ok(()) + } + + async fn apply_tun_update(&mut self, update: SharedVirtualNicMobileTunUpdate) { + self.drop_tun(); + self.tun_fd = update.tun_fd; + match self.open_tun().await { + Ok(()) => { + let _ = update.ack.send(Ok(())); + } + Err(err) => { + tracing::error!(fd = self.tun_fd, ?err, "failed to replace mobile tun"); + let _ = update.ack.send(Err(err)); + } + } + } + + async fn forward_member_packet_to_tun(&mut self, member_packet: SharedVirtualNicMemberPacket) { + let packet = self + .state + .prepare_member_packet_to_tun(member_packet.member_id, member_packet.packet); + let Some(tun_sink) = self.tun_sink.as_mut() else { + tracing::trace!( + member_id = ?member_packet.member_id, + "shared virtual nic dropped member packet without mobile tun" + ); + return; + }; + + if let Err(err) = tun_sink.send(packet).await { + tracing::error!(?err, "shared virtual nic write to mobile tun failed"); + self.drop_tun(); + } + } + + fn handle_control(&mut self, control: Option) -> bool { + match handle_dispatcher_control(&mut self.state, control) { + DispatcherControlResult::Continue => true, + DispatcherControlResult::Stop { invalidate, ack } => { + self.cleanup(invalidate); + acknowledge_dispatcher_shutdown(ack); + false + } + } + } + + fn drop_tun(&mut self) { + self.tun_stream.take(); + self.tun_sink.take(); + } + + fn cleanup(&mut self, invalidate: bool) { + self.drop_tun(); + if invalidate { + self.valid.store(false, Ordering::Release); + } + self.member_tunnel_table.detach_dispatcher(); + self.state.close_all(); + } +} + +#[cfg(any(mobile, test))] +async fn await_mobile_open(receiver: oneshot::Receiver>) -> Result<(), Error> { + receiver.await.map_err(|_| { + Error::from(anyhow::anyhow!( + "shared virtual nic mobile dispatcher stopped before opening TUN" + )) + })? +} + +#[derive(Default)] +struct SharedVirtualNicDispatcherState { + members: BTreeMap, + nat_table: SharedVirtualNicNatTable, + source_table: SharedVirtualNicSourceTable, +} + +impl SharedVirtualNicDispatcherState { + fn handle_control(&mut self, control: SharedVirtualNicControl) { + match control { + SharedVirtualNicControl::Register { member_id, entry } => { + self.register(member_id, entry); + } + SharedVirtualNicControl::Unregister { + member_id, + registration_id, + } => { + self.unregister(member_id, registration_id); + } + SharedVirtualNicControl::UpdateSources { + member_id, + sources, + ack, + } => { + self.nat_table.clear(); + self.source_table.update_member_sources(member_id, sources); + let _ = ack.send(()); + } + SharedVirtualNicControl::Shutdown { .. } => { + unreachable!("dispatcher shutdown is handled by the dispatcher task") + } + } + } + + fn register( + &mut self, + member_id: SharedVirtualNicMemberId, + entry: SharedVirtualNicMemberTunnelEntry, + ) { + let old_entry = self.members.insert(member_id, entry); + drop(old_entry); + } + + fn unregister( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) { + if self + .members + .get(&member_id) + .map(|entry| entry.registration_id) + != Some(registration_id) + { + return; + } + + let entry = self.members.remove(&member_id); + drop(entry); + self.nat_table.clear(); + self.source_table.remove_owner(member_id); + } + + fn close_all(&mut self) { + let members = std::mem::take(&mut self.members); + self.nat_table.clear(); + self.source_table.clear(); + + for entry in members.into_values() { + entry.close_notifier.notify_one(); + } + } + + fn prepare_member_packet_to_tun( + &mut self, + _member_id: SharedVirtualNicMemberId, + mut packet: ZCPacket, + ) -> ZCPacket { + self.nat_table.translate_reply(&mut packet); + packet + } + + async fn forward_tun_packet_to_member(&mut self, packet: ZCPacket) { + if !self.send_packet(packet).await { + tracing::trace!("shared virtual nic dropped packet without active member"); + } + } + + async fn send_packet(&mut self, packet: ZCPacket) -> bool { + let source_owner = self.source_table.owner_of_source(&packet, &self.members); + let preferred_destination_owner = source_owner.active_member(); + let destination_owner = self.source_table.owner_of_destination( + &packet, + &self.members, + preferred_destination_owner, + ); + if let Some(member_id) = destination_owner { + if should_translate_source_for_member(source_owner, member_id) { + return self + .send_packet_to_member_with_translation(member_id, packet) + .await + .is_ok(); + } else { + return self.send_packet_to_member(member_id, packet).await.is_ok(); + } + } + + match source_owner { + SourceOwner::Active(member_id) => { + return self.send_packet_to_member(member_id, packet).await.is_ok(); + } + SourceOwner::Inactive => return false, + SourceOwner::None => {} + } + + false + } + + async fn send_packet_to_member_with_translation( + &mut self, + member_id: SharedVirtualNicMemberId, + mut packet: ZCPacket, + ) -> Result<(), ZCPacket> { + let Some(key) = SharedVirtualNicFlowKey::from_packet(&packet) else { + return Err(packet); + }; + let Some(translated_src) = self + .source_table + .source_for_member_destination(member_id, key.dst) + else { + return self.send_packet_to_member(member_id, packet).await; + }; + let original_packet = packet.clone(); + let nat_entry = if key.src == translated_src { + None + } else if rewrite_packet_source(&mut packet, key.src, translated_src) { + Some((packet.clone(), key.src, translated_src)) + } else { + return Err(packet); + }; + + if self.send_packet_to_member(member_id, packet).await.is_err() { + return Err(original_packet); + } + if let Some((translated_packet, original_src, translated_src)) = nat_entry { + self.nat_table + .remember(&translated_packet, original_src, translated_src); + } + Ok(()) + } + + async fn send_packet_to_member( + &mut self, + member_id: SharedVirtualNicMemberId, + packet: ZCPacket, + ) -> Result<(), ZCPacket> { + let Some((registration_id, sender)) = self + .members + .get(&member_id) + .map(|entry| (entry.registration_id, entry.sender.clone())) + else { + return Err(packet); + }; + + match sender.send(packet).await { + Ok(()) => Ok(()), + Err(err) => { + self.unregister(member_id, registration_id); + Err(err.0) + } + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SourceOwner { + Active(SharedVirtualNicMemberId), + Inactive, + None, +} + +impl SourceOwner { + fn active_member(self) -> Option { + match self { + Self::Active(member_id) => Some(member_id), + Self::Inactive | Self::None => None, + } + } +} + +fn should_translate_source_for_member( + source_owner: SourceOwner, + member_id: SharedVirtualNicMemberId, +) -> bool { + !matches!(source_owner, SourceOwner::Active(source) if source == member_id) +} + +fn rewrite_packet_source( + packet: &mut ZCPacket, + expected_source: SharedVirtualNicFlowAddr, + new_source: SharedVirtualNicFlowAddr, +) -> bool { + rewrite_ip_addr(packet, expected_source, new_source, RewriteIpAddr::Source) +} + +fn rewrite_packet_destination( + packet: &mut ZCPacket, + expected_destination: SharedVirtualNicFlowAddr, + new_destination: SharedVirtualNicFlowAddr, +) -> bool { + rewrite_ip_addr( + packet, + expected_destination, + new_destination, + RewriteIpAddr::Destination, + ) +} + +#[derive(Clone, Copy)] +enum RewriteIpAddr { + Source, + Destination, +} + +fn rewrite_ip_addr( + packet: &mut ZCPacket, + expected_addr: SharedVirtualNicFlowAddr, + new_addr: SharedVirtualNicFlowAddr, + rewrite: RewriteIpAddr, +) -> bool { + match (expected_addr, new_addr) { + (SharedVirtualNicFlowAddr::V4(_), SharedVirtualNicFlowAddr::V4(_)) => { + rewrite_ipv4_addr(packet, expected_addr, new_addr, rewrite) + } + (SharedVirtualNicFlowAddr::V6(_), SharedVirtualNicFlowAddr::V6(_)) => { + rewrite_ipv6_addr(packet, expected_addr, new_addr, rewrite) + } + _ => false, + } +} + +fn rewrite_ipv4_addr( + packet: &mut ZCPacket, + expected_addr: SharedVirtualNicFlowAddr, + new_addr: SharedVirtualNicFlowAddr, + rewrite: RewriteIpAddr, +) -> bool { + let Some(expected_addr) = expected_addr.as_ipv4() else { + return false; + }; + let Some(new_addr) = new_addr.as_ipv4() else { + return false; + }; + + let payload = packet.mut_payload(); + let Some(mut ipv4_packet) = MutableIpv4Packet::new(payload) else { + return false; + }; + + let header_len = usize::from(ipv4_packet.get_header_length()) * 4; + if header_len < IPV4_HEADER_MIN_LEN || ipv4_packet.packet().len() < header_len { + return false; + } + + if ipv4_packet.get_fragment_offset() != 0 + || (ipv4_packet.get_flags() & ipv4::Ipv4Flags::MoreFragments) != 0 + { + return false; + } + + match rewrite { + RewriteIpAddr::Source if ipv4_packet.get_source() == expected_addr => { + ipv4_packet.set_source(new_addr); + } + RewriteIpAddr::Destination if ipv4_packet.get_destination() == expected_addr => { + ipv4_packet.set_destination(new_addr); + } + _ => return false, + } + + update_ipv4_transport_checksum(&mut ipv4_packet, header_len); + ipv4_packet.set_checksum(0); + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + true +} + +fn rewrite_ipv6_addr( + packet: &mut ZCPacket, + expected_addr: SharedVirtualNicFlowAddr, + new_addr: SharedVirtualNicFlowAddr, + rewrite: RewriteIpAddr, +) -> bool { + let Some(expected_addr) = expected_addr.as_ipv6() else { + return false; + }; + let Some(new_addr) = new_addr.as_ipv6() else { + return false; + }; + + let payload = packet.mut_payload(); + if payload.len() < IPV6_HEADER_LEN || payload[0] >> 4 != 6 { + return false; + } + let protocol = payload[6]; + if !matches!(protocol, TCP_PROTOCOL | UDP_PROTOCOL | ICMPV6_PROTOCOL) { + return false; + } + let Some(mut ipv6_packet) = MutableIpv6Packet::new(payload) else { + return false; + }; + + let old_source = ipv6_packet.get_source(); + let old_destination = ipv6_packet.get_destination(); + match rewrite { + RewriteIpAddr::Source if old_source == expected_addr => { + ipv6_packet.set_source(new_addr); + } + RewriteIpAddr::Destination if old_destination == expected_addr => { + ipv6_packet.set_destination(new_addr); + } + _ => return false, + } + + let source = ipv6_packet.get_source(); + let destination = ipv6_packet.get_destination(); + adjust_ipv6_transport_checksum( + ipv6_packet.packet_mut(), + protocol, + old_source, + source, + old_destination, + destination, + ); + true +} + +fn adjust_ipv6_transport_checksum( + payload: &mut [u8], + protocol: u8, + old_source: std::net::Ipv6Addr, + source: std::net::Ipv6Addr, + old_destination: std::net::Ipv6Addr, + destination: std::net::Ipv6Addr, +) { + let checksum_offset = match protocol { + TCP_PROTOCOL => 16, + UDP_PROTOCOL => 6, + ICMPV6_PROTOCOL => 2, + _ => return, + }; + let Some(checksum_bytes) = + payload.get_mut(IPV6_HEADER_LEN + checksum_offset..IPV6_HEADER_LEN + checksum_offset + 2) + else { + return; + }; + let checksum = u16::from_be_bytes([checksum_bytes[0], checksum_bytes[1]]); + let checksum = adjust_ipv6_pseudo_header_checksum( + checksum, + old_source, + source, + old_destination, + destination, + ); + let checksum = if protocol == UDP_PROTOCOL && checksum == 0 { + 0xffff + } else { + checksum + }; + checksum_bytes.copy_from_slice(&checksum.to_be_bytes()); +} + +fn update_ipv4_transport_checksum(ipv4_packet: &mut MutableIpv4Packet<'_>, header_len: usize) { + let source = ipv4_packet.get_source(); + let destination = ipv4_packet.get_destination(); + let protocol = ipv4_packet.get_next_level_protocol(); + let payload = ipv4_packet.packet_mut(); + let transport_payload = &mut payload[header_len..]; + + match protocol { + IpNextHeaderProtocols::Tcp => { + let Some(mut tcp_packet) = MutableTcpPacket::new(transport_payload) else { + return; + }; + tcp_packet.set_checksum(0); + let checksum = tcp::ipv4_checksum(&tcp_packet.to_immutable(), &source, &destination); + tcp_packet.set_checksum(checksum); + } + IpNextHeaderProtocols::Udp => { + let Some(mut udp_packet) = MutableUdpPacket::new(transport_payload) else { + return; + }; + if udp_packet.get_checksum() == 0 { + return; + } + udp_packet.set_checksum(0); + let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &source, &destination); + udp_packet.set_checksum(checksum); + } + IpNextHeaderProtocols::Icmp => { + let Some(mut icmp_packet) = MutableIcmpPacket::new(transport_payload) else { + return; + }; + icmp_packet.set_checksum(0); + let checksum = icmp::checksum(&icmp_packet.to_immutable()); + icmp_packet.set_checksum(checksum); + } + _ => {} + } +} + +fn adjust_ipv6_pseudo_header_checksum( + mut checksum: u16, + old_source: std::net::Ipv6Addr, + source: std::net::Ipv6Addr, + old_destination: std::net::Ipv6Addr, + destination: std::net::Ipv6Addr, +) -> u16 { + for (old, new) in old_source + .segments() + .into_iter() + .zip(source.segments()) + .chain( + old_destination + .segments() + .into_iter() + .zip(destination.segments()), + ) + { + checksum = adjust_checksum_word(checksum, old, new); + } + checksum +} + +fn adjust_checksum_word(checksum: u16, old_word: u16, new_word: u16) -> u16 { + let mut sum = u32::from(!checksum) + u32::from(!old_word) + u32::from(new_word); + while (sum >> 16) != 0 { + sum = (sum & 0xffff) + (sum >> 16); + } + !(sum as u16) +} + +#[derive(Default)] +struct SharedVirtualNicSourceTable { + member_sources: BTreeMap, +} + +impl SharedVirtualNicSourceTable { + fn update_member_sources( + &mut self, + member_id: SharedVirtualNicMemberId, + sources: SharedVirtualNicMemberSources, + ) { + if sources.is_empty() { + self.member_sources.remove(&member_id); + } else { + self.member_sources.insert(member_id, sources); + } + } + + fn remove_owner(&mut self, member_id: SharedVirtualNicMemberId) { + self.member_sources.remove(&member_id); + } + + fn clear(&mut self) { + self.member_sources.clear(); + } + + fn owner_of_source( + &self, + packet: &ZCPacket, + active_members: &BTreeMap, + ) -> SourceOwner { + let Some((source, _)) = SharedVirtualNicFlowAddr::packet_addrs(packet) else { + return SourceOwner::None; + }; + let mut found = false; + for (member_id, sources) in &self.member_sources { + if sources.owns_source(source) { + found = true; + if active_members.contains_key(member_id) { + return SourceOwner::Active(*member_id); + } + } + } + if found { + SourceOwner::Inactive + } else { + SourceOwner::None + } + } + + fn owner_of_destination( + &self, + packet: &ZCPacket, + active_members: &BTreeMap, + preferred_member: Option, + ) -> Option { + let (_, dst) = SharedVirtualNicFlowAddr::packet_addrs(packet)?; + + let mut first = None; + for (member_id, sources) in &self.member_sources { + if !active_members.contains_key(member_id) || !sources.has_destination(dst) { + continue; + } + if Some(*member_id) == preferred_member { + return Some(*member_id); + } + first.get_or_insert(*member_id); + } + first + } + + fn source_for_member_destination( + &self, + member_id: SharedVirtualNicMemberId, + dst: SharedVirtualNicFlowAddr, + ) -> Option { + let sources = self.member_sources.get(&member_id)?; + match dst { + SharedVirtualNicFlowAddr::V4(dst) => sources + .ipv4_source_for_destination(std::net::Ipv4Addr::from(dst)) + .map(SharedVirtualNicFlowAddr::from), + SharedVirtualNicFlowAddr::V6(dst) => sources + .ipv6_source_for_destination(std::net::Ipv6Addr::from(dst)) + .map(SharedVirtualNicFlowAddr::from), + } + } +} + +impl SharedVirtualNicMemberSources { + fn has_destination(&self, dst: SharedVirtualNicFlowAddr) -> bool { + match dst { + SharedVirtualNicFlowAddr::V4(dst) => { + let dst = std::net::Ipv4Addr::from(dst); + self.ipv4_addresses.iter().any(|addr| addr.contains(&dst)) + || self.ipv4_routes.iter().any(|route| route.contains(&dst)) + } + SharedVirtualNicFlowAddr::V6(dst) => { + let dst = std::net::Ipv6Addr::from(dst); + self.ipv6_addresses.iter().any(|addr| addr.contains(&dst)) + || self.ipv6_routes.iter().any(|route| route.contains(&dst)) + } + } + } + + fn ipv4_source_for_destination(&self, dst: std::net::Ipv4Addr) -> Option { + let mut best = None; + for addr in &self.ipv4_addresses { + if addr.contains(&dst) { + update_best_source(&mut best, addr.network_length(), addr.address()); + } + } + for route in &self.ipv4_routes { + if route.contains(&dst) { + let Some(source) = self.ipv4_source_for_route(route) else { + continue; + }; + update_best_source(&mut best, route.network_length(), source); + } + } + best.map(|(_, source)| source) + } + + fn ipv4_source_for_route(&self, route: &Ipv4Inet) -> Option { + let mut default_source = None; + for addr in &self.ipv4_addresses { + default_source.get_or_insert(addr.address()); + if route.contains(&addr.address()) { + return Some(addr.address()); + } + } + default_source + } + fn ipv6_source_for_destination(&self, dst: std::net::Ipv6Addr) -> Option { + let mut best = None; + for addr in &self.ipv6_addresses { + if addr.contains(&dst) { + update_best_source(&mut best, addr.network_length(), addr.address()); + } + } + for route in &self.ipv6_routes { + if route.contains(&dst) { + let Some(source) = self.ipv6_source_for_route(route) else { + continue; + }; + update_best_source(&mut best, route.network_length(), source); + } + } + best.map(|(_, source)| source) + } + + fn ipv6_source_for_route(&self, route: &Ipv6Inet) -> Option { + let mut default_source = None; + for addr in &self.ipv6_addresses { + default_source.get_or_insert(addr.address()); + if route.contains(&addr.address()) { + return Some(addr.address()); + } + } + default_source + } +} + +fn update_best_source(best: &mut Option<(u8, T)>, prefix: u8, source: T) { + if best + .map(|(best_prefix, _)| prefix > best_prefix) + .unwrap_or(true) + { + *best = Some((prefix, source)); + } +} + +fn transport_ports(protocol: u8, payload: &[u8]) -> Option { + let min_len = match protocol { + TCP_PROTOCOL => TCP_HEADER_MIN_LEN, + UDP_PROTOCOL => UDP_HEADER_LEN, + ICMP_PROTOCOL | ICMPV6_PROTOCOL => ICMP_ECHO_HEADER_LEN, + _ => return None, + }; + + if payload.len() < min_len { + return None; + } + + match protocol { + ICMP_PROTOCOL => icmp_echo_flow(payload), + ICMPV6_PROTOCOL => icmpv6_echo_flow(payload), + _ => Some(SharedVirtualNicTransportPorts { + src: u16::from_be_bytes([payload[0], payload[1]]), + dst: u16::from_be_bytes([payload[2], payload[3]]), + }), + } +} + +fn icmp_echo_flow(payload: &[u8]) -> Option { + match payload[0] { + ty if ty == icmp::IcmpTypes::EchoRequest.0 || ty == icmp::IcmpTypes::EchoReply.0 => { + Some(SharedVirtualNicTransportPorts { + src: u16::from_be_bytes([payload[4], payload[5]]), + dst: u16::from_be_bytes([payload[6], payload[7]]), + }) + } + _ => None, + } +} + +fn icmpv6_echo_flow(payload: &[u8]) -> Option { + match payload[0] { + 128 | 129 => Some(SharedVirtualNicTransportPorts { + src: u16::from_be_bytes([payload[4], payload[5]]), + dst: u16::from_be_bytes([payload[6], payload[7]]), + }), + _ => None, + } +} + +fn read_ipv6_addr(payload: &[u8], start: usize) -> [u8; 16] { + let mut addr = [0; 16]; + addr.copy_from_slice(&payload[start..start + 16]); + addr +} + +#[cfg(test)] +mod tests { + use std::{ + net::{Ipv4Addr, Ipv6Addr}, + time::Duration, + }; + + use super::*; + use easytier_core::tunnel::{ + TunnelError, ring::create_ring_tunnel_pair, wrapper::TunnelWrapper, + }; + + #[tokio::test] + async fn mobile_open_handshake_waits_for_dispatcher_result() { + let (sender, receiver) = oneshot::channel(); + let waiter = tokio::spawn(await_mobile_open(receiver)); + + tokio::task::yield_now().await; + assert!(!waiter.is_finished()); + + sender.send(Ok(())).unwrap(); + waiter.await.unwrap().unwrap(); + } + + #[tokio::test] + async fn mobile_open_handshake_preserves_original_error() { + let (sender, receiver) = oneshot::channel(); + let waiter = tokio::spawn(await_mobile_open(receiver)); + let errno = 9; + + sender + .send(Err(std::io::Error::from_raw_os_error(errno).into())) + .unwrap(); + let err = waiter.await.unwrap().unwrap_err(); + assert!(matches!(err, Error::IOError(ref io_err) if io_err.raw_os_error() == Some(errno))); + } + + fn ipv6_packet(src: Ipv6Addr, dst: Ipv6Addr) -> ZCPacket { + let mut payload = vec![0; IPV6_HEADER_LEN]; + payload[0] = 0x60; + payload[6] = 58; + payload[8..24].copy_from_slice(&src.octets()); + payload[24..40].copy_from_slice(&dst.octets()); + ZCPacket::new_with_payload(&payload) + } + + fn ipv6_udp_packet_with_ports( + src: Ipv6Addr, + dst: Ipv6Addr, + src_port: u16, + dst_port: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV6_HEADER_LEN + UDP_HEADER_LEN]; + { + let mut ipv6_packet = pnet_packet::ipv6::MutableIpv6Packet::new(&mut payload).unwrap(); + ipv6_packet.set_version(6); + ipv6_packet.set_payload_length(UDP_HEADER_LEN as u16); + ipv6_packet.set_next_header(IpNextHeaderProtocols::Udp); + ipv6_packet.set_hop_limit(64); + ipv6_packet.set_source(src); + ipv6_packet.set_destination(dst); + } + { + let mut udp_packet = MutableUdpPacket::new(&mut payload[IPV6_HEADER_LEN..]).unwrap(); + udp_packet.set_source(src_port); + udp_packet.set_destination(dst_port); + udp_packet.set_length(UDP_HEADER_LEN as u16); + let checksum = udp::ipv6_checksum(&udp_packet.to_immutable(), &src, &dst); + udp_packet.set_checksum(checksum); + } + ZCPacket::new_with_payload(&payload) + } + + fn ipv6_tcp_packet_with_ports( + src: Ipv6Addr, + dst: Ipv6Addr, + src_port: u16, + dst_port: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV6_HEADER_LEN + TCP_HEADER_MIN_LEN]; + { + let mut ipv6_packet = pnet_packet::ipv6::MutableIpv6Packet::new(&mut payload).unwrap(); + ipv6_packet.set_version(6); + ipv6_packet.set_payload_length(TCP_HEADER_MIN_LEN as u16); + ipv6_packet.set_next_header(IpNextHeaderProtocols::Tcp); + ipv6_packet.set_hop_limit(64); + ipv6_packet.set_source(src); + ipv6_packet.set_destination(dst); + } + { + let mut tcp_packet = MutableTcpPacket::new(&mut payload[IPV6_HEADER_LEN..]).unwrap(); + tcp_packet.set_source(src_port); + tcp_packet.set_destination(dst_port); + tcp_packet.set_data_offset(5); + let checksum = tcp::ipv6_checksum(&tcp_packet.to_immutable(), &src, &dst); + tcp_packet.set_checksum(checksum); + } + ZCPacket::new_with_payload(&payload) + } + + fn ipv6_icmp_echo_packet(src: Ipv6Addr, dst: Ipv6Addr) -> ZCPacket { + ipv6_icmp_packet_with_id(src, dst, false, 1234) + } + + fn ipv6_icmp_packet_with_id( + src: Ipv6Addr, + dst: Ipv6Addr, + reply: bool, + identifier: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV6_HEADER_LEN + ICMP_ECHO_HEADER_LEN]; + { + let mut ipv6_packet = pnet_packet::ipv6::MutableIpv6Packet::new(&mut payload).unwrap(); + ipv6_packet.set_version(6); + ipv6_packet.set_payload_length(ICMP_ECHO_HEADER_LEN as u16); + ipv6_packet.set_next_header(IpNextHeaderProtocols::Icmpv6); + ipv6_packet.set_hop_limit(64); + ipv6_packet.set_source(src); + ipv6_packet.set_destination(dst); + } + { + let mut icmp_packet = + pnet_packet::icmpv6::MutableIcmpv6Packet::new(&mut payload[IPV6_HEADER_LEN..]) + .unwrap(); + icmp_packet.set_icmpv6_type(if reply { + pnet_packet::icmpv6::Icmpv6Types::EchoReply + } else { + pnet_packet::icmpv6::Icmpv6Types::EchoRequest + }); + icmp_packet.packet_mut()[4..6].copy_from_slice(&identifier.to_be_bytes()); + icmp_packet.packet_mut()[6..8].copy_from_slice(&1_u16.to_be_bytes()); + let checksum = pnet_packet::icmpv6::checksum(&icmp_packet.to_immutable(), &src, &dst); + icmp_packet.set_checksum(checksum); + } + ZCPacket::new_with_payload(&payload) + } + + fn ipv4_udp_packet(src: Ipv4Addr, dst: Ipv4Addr) -> ZCPacket { + ipv4_udp_packet_with_ports(src, dst, 1234, 5678) + } + + fn ipv4_icmp_echo_packet(src: Ipv4Addr, dst: Ipv4Addr) -> ZCPacket { + ipv4_icmp_packet_with_id(src, dst, icmp::IcmpTypes::EchoRequest, 0, 0) + } + + fn ipv4_icmp_packet_with_id( + src: Ipv4Addr, + dst: Ipv4Addr, + icmp_type: icmp::IcmpType, + identifier: u16, + sequence: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV4_HEADER_MIN_LEN + 8]; + let payload_len = payload.len(); + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut payload).unwrap(); + ipv4_packet.set_version(4); + ipv4_packet.set_header_length(5); + ipv4_packet.set_total_length(payload_len as u16); + ipv4_packet.set_ttl(64); + ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Icmp); + ipv4_packet.set_source(src); + ipv4_packet.set_destination(dst); + } + { + let mut icmp_packet = + MutableIcmpPacket::new(&mut payload[IPV4_HEADER_MIN_LEN..]).unwrap(); + icmp_packet.set_icmp_type(icmp_type); + icmp_packet.set_icmp_code(icmp::IcmpCode(0)); + icmp_packet.packet_mut()[4..6].copy_from_slice(&identifier.to_be_bytes()); + icmp_packet.packet_mut()[6..8].copy_from_slice(&sequence.to_be_bytes()); + let checksum = icmp::checksum(&icmp_packet.to_immutable()); + icmp_packet.set_checksum(checksum); + } + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut payload).unwrap(); + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + } + ZCPacket::new_with_payload(&payload) + } + + fn ipv4_udp_packet_with_ports( + src: Ipv4Addr, + dst: Ipv4Addr, + src_port: u16, + dst_port: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV4_HEADER_MIN_LEN + UDP_HEADER_LEN]; + let payload_len = payload.len(); + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut payload).unwrap(); + ipv4_packet.set_version(4); + ipv4_packet.set_header_length(5); + ipv4_packet.set_total_length(payload_len as u16); + ipv4_packet.set_ttl(64); + ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp); + ipv4_packet.set_source(src); + ipv4_packet.set_destination(dst); + } + { + let mut udp_packet = + MutableUdpPacket::new(&mut payload[IPV4_HEADER_MIN_LEN..]).unwrap(); + udp_packet.set_source(src_port); + udp_packet.set_destination(dst_port); + udp_packet.set_length(UDP_HEADER_LEN as u16); + let checksum = udp::ipv4_checksum(&udp_packet.to_immutable(), &src, &dst); + udp_packet.set_checksum(checksum); + } + { + let mut ipv4_packet = MutableIpv4Packet::new(&mut payload).unwrap(); + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + } + ZCPacket::new_with_payload(&payload) + } + + fn member_sources(ipv4: &[&str]) -> SharedVirtualNicMemberSources { + SharedVirtualNicMemberSources::from_claims(&member_claims(ipv4, &[], &[], &[])) + } + + fn member_sources_with_ipv4_routes( + ipv4: &[&str], + ipv4_routes: &[&str], + ) -> SharedVirtualNicMemberSources { + SharedVirtualNicMemberSources::from_claims(&member_claims(ipv4, &[], ipv4_routes, &[])) + } + + fn member_sources_with_ipv6(ipv6: &[&str]) -> SharedVirtualNicMemberSources { + SharedVirtualNicMemberSources::from_claims(&member_claims(&[], ipv6, &[], &[])) + } + + fn member_sources_with_ipv6_routes( + ipv6: &[&str], + ipv6_routes: &[&str], + ) -> SharedVirtualNicMemberSources { + SharedVirtualNicMemberSources::from_claims(&member_claims(&[], ipv6, &[], ipv6_routes)) + } + + fn member_claims( + ipv4: &[&str], + ipv6: &[&str], + ipv4_routes: &[&str], + ipv6_routes: &[&str], + ) -> SharedIfConfigClaims { + SharedIfConfigClaims { + ipv4_addresses: ipv4.iter().map(|addr| addr.parse().unwrap()).collect(), + ipv6_addresses: ipv6.iter().map(|addr| addr.parse().unwrap()).collect(), + ipv4_routes: ipv4_routes + .iter() + .map(|route| { + let inet = route.parse::().unwrap(); + SharedIpv4Route::new(inet.address(), inet.network_length(), None) + }) + .collect(), + ipv6_routes: ipv6_routes + .iter() + .map(|route| { + let inet = route.parse::().unwrap(); + SharedIpv6Route::new(inet.address(), inet.network_length(), None) + }) + .collect(), + mtu: None, + } + } + + fn member_entry(sender: mpsc::Sender) -> SharedVirtualNicMemberTunnelEntry { + member_entry_with_registration(sender, uuid::Uuid::from_u128(1)) + } + + fn member_entry_with_registration( + sender: mpsc::Sender, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> SharedVirtualNicMemberTunnelEntry { + SharedVirtualNicMemberTunnelEntry { + registration_id, + sender, + close_notifier: Arc::new(Notify::new()), + _tasks: Vec::new(), + } + } + + #[test] + fn source_table_selects_ipv6_source_owner() { + let first = uuid::Uuid::from_u128(1); + let second = uuid::Uuid::from_u128(2); + let second_addr = "2001:db8::2".parse::().unwrap(); + let dst = "2001:db8:ffff::1".parse::().unwrap(); + let mut table = SharedVirtualNicSourceTable::default(); + let (first_sender, _first_receiver) = mpsc::channel(1); + let (second_sender, _second_receiver) = mpsc::channel(1); + let mut members = BTreeMap::new(); + members.insert(first, member_entry(first_sender)); + members.insert(second, member_entry(second_sender)); + + table.update_member_sources(first, member_sources_with_ipv6(&["2001:db8::1/64"])); + table.update_member_sources(second, member_sources_with_ipv6(&["2001:db8::2/64"])); + + assert_eq!( + table.owner_of_source(&ipv6_packet(second_addr, dst), &members), + SourceOwner::Active(second) + ); + + table.remove_owner(second); + assert_eq!( + table.owner_of_source(&ipv6_packet(second_addr, dst), &members), + SourceOwner::None + ); + } + + #[test] + fn source_table_selects_ipv4_destination_owner_from_route_claim() { + let first = uuid::Uuid::from_u128(1); + let second = uuid::Uuid::from_u128(2); + let src = Ipv4Addr::new(100, 64, 0, 1); + let dst = Ipv4Addr::new(10, 99, 0, 2); + let mut table = SharedVirtualNicSourceTable::default(); + let (first_sender, _first_receiver) = mpsc::channel(1); + let (second_sender, _second_receiver) = mpsc::channel(1); + let mut members = BTreeMap::new(); + members.insert(first, member_entry(first_sender)); + members.insert(second, member_entry(second_sender)); + + table.update_member_sources(first, member_sources(&["10.231.1.1/24"])); + table.update_member_sources( + second, + member_sources_with_ipv4_routes(&["10.231.2.1/24"], &["10.99.0.0/24"]), + ); + + assert_eq!( + table.owner_of_destination(&ipv4_udp_packet(src, dst), &members, None), + Some(second) + ); + } + + #[tokio::test] + async fn dispatcher_prefers_source_owner_for_shared_magic_dns() { + let first = uuid::Uuid::from_u128(1); + let source_owner = uuid::Uuid::from_u128(2); + let first_ip = Ipv4Addr::new(10, 231, 1, 1); + let source_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(100, 100, 100, 101); + let (first_sender, mut first_receiver) = mpsc::channel(1); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(first, member_entry(first_sender)); + state.register(source_owner, member_entry(source_sender)); + state.source_table.update_member_sources( + first, + member_sources_with_ipv4_routes(&["10.231.1.1/24"], &["100.100.100.101/32"]), + ); + state.source_table.update_member_sources( + source_owner, + member_sources_with_ipv4_routes(&["10.231.2.1/24"], &["100.100.100.101/32"]), + ); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(source_owner_ip, remote_ip)) + .await; + + assert!(first_receiver.try_recv().is_err()); + let packet = source_receiver.try_recv().unwrap(); + let ipv4 = pnet_packet::ipv4::Ipv4Packet::new(packet.payload()).unwrap(); + assert_eq!(ipv4.get_source(), source_owner_ip); + assert_ne!(ipv4.get_source(), first_ip); + assert_eq!(ipv4.get_destination(), remote_ip); + } + + #[test] + fn rewrite_ipv4_source_rejects_fragments_without_changing_packet() { + let src = Ipv4Addr::new(10, 231, 1, 1); + let dst = Ipv4Addr::new(10, 231, 2, 2); + for (flags, offset) in [(ipv4::Ipv4Flags::MoreFragments, 0), (0, 1)] { + let mut packet = ipv4_udp_packet(src, dst); + let mut ipv4 = MutableIpv4Packet::new(packet.mut_payload()).unwrap(); + ipv4.set_flags(flags); + ipv4.set_fragment_offset(offset); + let original = packet.payload().to_vec(); + + assert!(!rewrite_packet_source( + &mut packet, + SharedVirtualNicFlowAddr::from(src), + SharedVirtualNicFlowAddr::from(Ipv4Addr::new(10, 231, 2, 1)) + )); + assert_eq!(packet.payload(), original); + } + } + + #[tokio::test] + async fn dispatcher_prefers_source_owner_over_first_member() { + let first = uuid::Uuid::from_u128(1); + let owner = uuid::Uuid::from_u128(2); + let source = "2001:db8::2".parse::().unwrap(); + let dst = "2001:db8:ffff::1".parse::().unwrap(); + let (first_sender, mut first_receiver) = mpsc::channel(1); + let (owner_sender, mut owner_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(first, member_entry(first_sender)); + state.register(owner, member_entry(owner_sender)); + state + .source_table + .update_member_sources(owner, member_sources_with_ipv6(&["2001:db8::2/64"])); + state + .forward_tun_packet_to_member(ipv6_packet(source, dst)) + .await; + + assert!(first_receiver.try_recv().is_err()); + assert!(owner_receiver.try_recv().is_ok()); + } + + #[tokio::test] + async fn dispatcher_forwards_external_ipv6_to_route_only_owner() { + let route_owner = uuid::Uuid::from_u128(1); + let source = "2001:db8:ffff::2".parse::().unwrap(); + let destination = "2001:db8:100::2".parse::().unwrap(); + let (sender, mut receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(route_owner, member_entry(sender)); + state.source_table.update_member_sources( + route_owner, + member_sources_with_ipv6_routes(&[], &["2001:db8:100::2/128"]), + ); + state + .forward_tun_packet_to_member(ipv6_packet(source, destination)) + .await; + + let packet = receiver.try_recv().unwrap(); + let ipv6 = pnet_packet::ipv6::Ipv6Packet::new(packet.payload()).unwrap(); + assert_eq!(ipv6.get_source(), source); + assert_eq!(ipv6.get_destination(), destination); + } + + #[tokio::test] + async fn dispatcher_drops_inactive_source_owner_without_fallback() { + let fallback = uuid::Uuid::from_u128(1); + let owner = uuid::Uuid::from_u128(2); + let source = "2001:db8::2".parse::().unwrap(); + let dst = "2001:db8:ffff::1".parse::().unwrap(); + let (fallback_sender, mut fallback_receiver) = mpsc::channel(1); + let (owner_sender, _owner_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(fallback, member_entry(fallback_sender)); + state.register(owner, member_entry(owner_sender)); + state + .source_table + .update_member_sources(owner, member_sources_with_ipv6(&["2001:db8::2/64"])); + state.unregister(owner, uuid::Uuid::from_u128(1)); + state + .forward_tun_packet_to_member(ipv6_packet(source, dst)) + .await; + + assert!(fallback_receiver.try_recv().is_err()); + } + + #[tokio::test] + async fn dispatcher_drops_unknown_source_and_destination_without_fallback() { + let member_id = uuid::Uuid::from_u128(1); + let unknown_source = Ipv4Addr::new(100, 64, 0, 1); + let unknown_destination = Ipv4Addr::new(203, 0, 113, 1); + let (sender, mut receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(member_id, member_entry(sender)); + state + .source_table + .update_member_sources(member_id, member_sources(&["10.231.1.1/24"])); + state + .forward_tun_packet_to_member(ipv4_udp_packet(unknown_source, unknown_destination)) + .await; + + assert!(receiver.try_recv().is_err()); + } + + #[tokio::test] + async fn dispatcher_translates_wrong_local_ipv4_source_to_destination_member() { + let source_owner = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let source_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(10, 231, 2, 2); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(source_owner, member_entry(source_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(source_owner, member_sources(&["10.231.1.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.2.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(source_owner_ip, remote_ip)) + .await; + + assert!(source_receiver.try_recv().is_err()); + let translated = destination_receiver.try_recv().unwrap(); + let translated_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(translated.payload()).unwrap(); + assert_eq!(translated_ipv4.get_source(), destination_owner_ip); + assert_eq!(translated_ipv4.get_destination(), remote_ip); + + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_udp_packet_with_ports(remote_ip, destination_owner_ip, 5678, 1234), + ); + let reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv4.get_source(), remote_ip); + assert_eq!(reply_ipv4.get_destination(), source_owner_ip); + } + + #[tokio::test] + async fn dispatcher_translates_wrong_local_ipv6_source_to_destination_member() { + let source_owner = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let source_owner_ip = "2001:db8:1::1".parse::().unwrap(); + let destination_owner_ip = "2001:db8:2::1".parse::().unwrap(); + let remote_ip = "2001:db8:99::2".parse::().unwrap(); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(source_owner, member_entry(source_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state.source_table.update_member_sources( + source_owner, + member_sources_with_ipv6(&["2001:db8:1::1/64"]), + ); + state.source_table.update_member_sources( + destination_owner, + member_sources_with_ipv6_routes(&["2001:db8:2::1/64"], &["2001:db8:99::/64"]), + ); + + state + .forward_tun_packet_to_member(ipv6_udp_packet_with_ports( + source_owner_ip, + remote_ip, + 1234, + 5678, + )) + .await; + + assert!(source_receiver.try_recv().is_err()); + let translated = destination_receiver.try_recv().unwrap(); + let translated_ipv6 = pnet_packet::ipv6::Ipv6Packet::new(translated.payload()).unwrap(); + assert_eq!(translated_ipv6.get_source(), destination_owner_ip); + assert_eq!(translated_ipv6.get_destination(), remote_ip); + let translated_udp = + pnet_packet::udp::UdpPacket::new(&translated.payload()[IPV6_HEADER_LEN..]).unwrap(); + assert_eq!( + translated_udp.get_checksum(), + udp::ipv6_checksum(&translated_udp, &destination_owner_ip, &remote_ip,) + ); + + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv6_udp_packet_with_ports(remote_ip, destination_owner_ip, 5678, 1234), + ); + let reply_ipv6 = pnet_packet::ipv6::Ipv6Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv6.get_source(), remote_ip); + assert_eq!(reply_ipv6.get_destination(), source_owner_ip); + let reply_udp = + pnet_packet::udp::UdpPacket::new(&reply.payload()[IPV6_HEADER_LEN..]).unwrap(); + assert_eq!( + reply_udp.get_checksum(), + udp::ipv6_checksum(&reply_udp, &remote_ip, &source_owner_ip) + ); + } + + #[test] + fn rewrite_ipv6_source_updates_tcp_and_icmpv6_checksums() { + let source = "2001:db8:1::1".parse::().unwrap(); + let translated_source = "2001:db8:2::1".parse::().unwrap(); + let destination = "2001:db8:99::2".parse::().unwrap(); + + let mut tcp_packet = ipv6_tcp_packet_with_ports(source, destination, 1234, 5678); + assert!(rewrite_packet_source( + &mut tcp_packet, + SharedVirtualNicFlowAddr::from(source), + SharedVirtualNicFlowAddr::from(translated_source), + )); + let tcp = + pnet_packet::tcp::TcpPacket::new(&tcp_packet.payload()[IPV6_HEADER_LEN..]).unwrap(); + assert_eq!( + tcp.get_checksum(), + tcp::ipv6_checksum(&tcp, &translated_source, &destination) + ); + + let mut icmp_packet = ipv6_icmp_echo_packet(source, destination); + assert!(rewrite_packet_source( + &mut icmp_packet, + SharedVirtualNicFlowAddr::from(source), + SharedVirtualNicFlowAddr::from(translated_source), + )); + let icmp = + pnet_packet::icmpv6::Icmpv6Packet::new(&icmp_packet.payload()[IPV6_HEADER_LEN..]) + .unwrap(); + assert_eq!( + icmp.get_checksum(), + pnet_packet::icmpv6::checksum(&icmp, &translated_source, &destination) + ); + } + + #[test] + fn icmpv6_nat_replies_keep_distinct_source_networks() { + let first_source = "2001:db8:1::1".parse::().unwrap(); + let second_source = "2001:db8:2::1".parse::().unwrap(); + let translated_source = "2001:db8:3::1".parse::().unwrap(); + let destination = "2001:db8:3::2".parse::().unwrap(); + let mut nat = SharedVirtualNicNatTable::default(); + + for (identifier, source) in [(100, first_source), (200, second_source)] { + nat.remember( + &ipv6_icmp_packet_with_id(translated_source, destination, false, identifier), + source.into(), + translated_source.into(), + ); + } + + for (identifier, source) in [(100, first_source), (200, second_source)] { + let mut reply = + ipv6_icmp_packet_with_id(destination, translated_source, true, identifier); + assert!(nat.translate_reply(&mut reply)); + let ipv6 = pnet_packet::ipv6::Ipv6Packet::new(reply.payload()).unwrap(); + assert_eq!(ipv6.get_destination(), source); + } + } + + #[test] + fn rewrite_ipv6_source_rejects_extension_headers_without_changing_packet() { + let source = "2001:db8:1::1".parse::().unwrap(); + let destination = "2001:db8:99::2".parse::().unwrap(); + let mut packet = ipv6_packet(source, destination); + packet.mut_payload()[6] = 0; // Hop-by-hop options header. + let original = packet.payload().to_vec(); + + assert!(!rewrite_packet_source( + &mut packet, + SharedVirtualNicFlowAddr::from(source), + SharedVirtualNicFlowAddr::from("2001:db8:2::1".parse::().unwrap()), + )); + assert_eq!(packet.payload(), original); + } + + #[tokio::test] + async fn dispatcher_unregister_clears_nat_translation() { + let source_owner = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let source_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(10, 231, 2, 2); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(source_owner, member_entry(source_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(source_owner, member_sources(&["10.231.1.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.2.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(source_owner_ip, remote_ip)) + .await; + + assert!(source_receiver.try_recv().is_err()); + assert!(destination_receiver.try_recv().is_ok()); + + state.unregister(destination_owner, uuid::Uuid::from_u128(1)); + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_udp_packet_with_ports(remote_ip, destination_owner_ip, 5678, 1234), + ); + let reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv4.get_source(), remote_ip); + assert_eq!(reply_ipv4.get_destination(), destination_owner_ip); + } + + #[tokio::test] + async fn dispatcher_unregister_source_owner_clears_nat_translation() { + let source_owner = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let source_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(10, 231, 2, 2); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(source_owner, member_entry(source_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(source_owner, member_sources(&["10.231.1.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.2.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(source_owner_ip, remote_ip)) + .await; + + assert!(source_receiver.try_recv().is_err()); + assert!(destination_receiver.try_recv().is_ok()); + + state.unregister(source_owner, uuid::Uuid::from_u128(1)); + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_udp_packet_with_ports(remote_ip, destination_owner_ip, 5678, 1234), + ); + let reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv4.get_source(), remote_ip); + assert_eq!(reply_ipv4.get_destination(), destination_owner_ip); + } + + #[tokio::test] + async fn dispatcher_update_sources_clears_nat_translation() { + let source_owner = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let source_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(10, 231, 2, 2); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(source_owner, member_entry(source_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(source_owner, member_sources(&["10.231.1.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.2.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(source_owner_ip, remote_ip)) + .await; + + assert!(source_receiver.try_recv().is_err()); + assert!(destination_receiver.try_recv().is_ok()); + + let (ack, _rx) = oneshot::channel(); + state.handle_control(SharedVirtualNicControl::UpdateSources { + member_id: destination_owner, + sources: member_sources(&["10.231.3.1/24"]), + ack, + }); + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_udp_packet_with_ports(remote_ip, destination_owner_ip, 5678, 1234), + ); + let reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv4.get_source(), remote_ip); + assert_eq!(reply_ipv4.get_destination(), destination_owner_ip); + } + + #[tokio::test] + async fn dispatcher_translates_unknown_ipv4_source_to_route_owner() { + let first = uuid::Uuid::from_u128(1); + let route_owner = uuid::Uuid::from_u128(2); + let synthetic_source = Ipv4Addr::new(100, 64, 0, 1); + let route_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let remote_ip = Ipv4Addr::new(10, 99, 0, 2); + let (first_sender, mut first_receiver) = mpsc::channel(1); + let (route_sender, mut route_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(first, member_entry(first_sender)); + state.register(route_owner, member_entry(route_sender)); + state + .source_table + .update_member_sources(first, member_sources(&["10.231.1.1/24"])); + state.source_table.update_member_sources( + route_owner, + member_sources_with_ipv4_routes(&["10.231.2.1/24"], &["10.99.0.0/24"]), + ); + + state + .forward_tun_packet_to_member(ipv4_udp_packet(synthetic_source, remote_ip)) + .await; + + assert!(first_receiver.try_recv().is_err()); + let translated = route_receiver.try_recv().unwrap(); + let translated_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(translated.payload()).unwrap(); + assert_eq!(translated_ipv4.get_source(), route_owner_ip); + assert_eq!(translated_ipv4.get_destination(), remote_ip); + } + + #[tokio::test] + async fn dispatcher_translates_android_synthetic_icmp_source_to_destination_member() { + let first = uuid::Uuid::from_u128(1); + let destination_owner = uuid::Uuid::from_u128(2); + let synthetic_source = Ipv4Addr::new(100, 64, 0, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let remote_ip = Ipv4Addr::new(10, 231, 1, 2); + let (first_sender, mut first_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(first, member_entry(first_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(first, member_sources(&["10.231.2.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.1.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_icmp_echo_packet(synthetic_source, remote_ip)) + .await; + + assert!(first_receiver.try_recv().is_err()); + let translated = destination_receiver.try_recv().unwrap(); + let translated_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(translated.payload()).unwrap(); + assert_eq!(translated_ipv4.get_source(), destination_owner_ip); + assert_eq!(translated_ipv4.get_destination(), remote_ip); + + let reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_icmp_echo_packet(remote_ip, destination_owner_ip), + ); + let reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(reply.payload()).unwrap(); + assert_eq!(reply_ipv4.get_source(), remote_ip); + assert_eq!(reply_ipv4.get_destination(), synthetic_source); + } + + #[tokio::test] + async fn dispatcher_keeps_distinct_icmp_nat_entries_by_echo_id() { + let first_source = uuid::Uuid::from_u128(1); + let second_source = uuid::Uuid::from_u128(2); + let destination_owner = uuid::Uuid::from_u128(3); + let first_source_ip = Ipv4Addr::new(10, 231, 1, 1); + let second_source_ip = Ipv4Addr::new(10, 231, 2, 1); + let destination_owner_ip = Ipv4Addr::new(10, 231, 3, 1); + let remote_ip = Ipv4Addr::new(10, 231, 3, 2); + let (first_sender, mut first_receiver) = mpsc::channel(1); + let (second_sender, mut second_receiver) = mpsc::channel(1); + let (destination_sender, mut destination_receiver) = mpsc::channel(2); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(first_source, member_entry(first_sender)); + state.register(second_source, member_entry(second_sender)); + state.register(destination_owner, member_entry(destination_sender)); + state + .source_table + .update_member_sources(first_source, member_sources(&["10.231.1.1/24"])); + state + .source_table + .update_member_sources(second_source, member_sources(&["10.231.2.1/24"])); + state + .source_table + .update_member_sources(destination_owner, member_sources(&["10.231.3.1/24"])); + + state + .forward_tun_packet_to_member(ipv4_icmp_packet_with_id( + first_source_ip, + remote_ip, + icmp::IcmpTypes::EchoRequest, + 100, + 1, + )) + .await; + state + .forward_tun_packet_to_member(ipv4_icmp_packet_with_id( + second_source_ip, + remote_ip, + icmp::IcmpTypes::EchoRequest, + 200, + 1, + )) + .await; + + assert!(first_receiver.try_recv().is_err()); + assert!(second_receiver.try_recv().is_err()); + assert!(destination_receiver.try_recv().is_ok()); + assert!(destination_receiver.try_recv().is_ok()); + + let first_reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_icmp_packet_with_id( + remote_ip, + destination_owner_ip, + icmp::IcmpTypes::EchoReply, + 100, + 1, + ), + ); + let first_reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(first_reply.payload()).unwrap(); + assert_eq!(first_reply_ipv4.get_destination(), first_source_ip); + + let second_reply = state.prepare_member_packet_to_tun( + destination_owner, + ipv4_icmp_packet_with_id( + remote_ip, + destination_owner_ip, + icmp::IcmpTypes::EchoReply, + 200, + 1, + ), + ); + let second_reply_ipv4 = pnet_packet::ipv4::Ipv4Packet::new(second_reply.payload()).unwrap(); + assert_eq!(second_reply_ipv4.get_destination(), second_source_ip); + } + + #[tokio::test] + async fn dispatcher_preserves_owned_ipv4_source_with_unspecified_placeholder() { + let member_id = uuid::Uuid::from_u128(1); + let local_ip = Ipv4Addr::new(10, 231, 1, 1); + let remote_ip = Ipv4Addr::new(10, 231, 1, 2); + let (sender, mut receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(member_id, member_entry(sender)); + state + .source_table + .update_member_sources(member_id, member_sources(&["0.0.0.0/0", "10.231.1.1/24"])); + state + .forward_tun_packet_to_member(ipv4_udp_packet(local_ip, remote_ip)) + .await; + + let packet = receiver.try_recv().unwrap(); + let ipv4 = pnet_packet::ipv4::Ipv4Packet::new(packet.payload()).unwrap(); + assert_eq!(ipv4.get_source(), local_ip); + assert_eq!(ipv4.get_destination(), remote_ip); + } + + #[tokio::test] + async fn dispatcher_ignores_stale_member_unregister() { + let member_id = uuid::Uuid::from_u128(1); + let stale_registration = uuid::Uuid::from_u128(10); + let current_registration = uuid::Uuid::from_u128(11); + let src = "2001:db8::1".parse::().unwrap(); + let dst = "2001:db8:ffff::1".parse::().unwrap(); + let (sender, mut receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register( + member_id, + member_entry_with_registration(sender, current_registration), + ); + state + .source_table + .update_member_sources(member_id, member_sources_with_ipv6(&["2001:db8::1/64"])); + state.unregister(member_id, stale_registration); + state + .forward_tun_packet_to_member(ipv6_packet(src, dst)) + .await; + + assert!(receiver.try_recv().is_ok()); + } + + #[tokio::test] + async fn dispatcher_invalidates_shared_nic_when_tun_read_fails() { + let member_id = uuid::Uuid::from_u128(1); + let (tun_tx, tun_rx) = mpsc::unbounded_channel(); + let tun_stream = futures::stream::unfold(tun_rx, |mut receiver| async move { + receiver.recv().await.map(|packet| (packet, receiver)) + }); + let tun_sink = futures::sink::unfold((), |(), _packet: ZCPacket| async { + Ok::<(), TunnelError>(()) + }); + let tunnel = TunnelWrapper::new(tun_stream, tun_sink, None); + let member_tunnel_table = SharedVirtualNicMemberTunnelTable::default(); + let valid = Arc::new(AtomicBool::new(true)); + let dispatcher = SharedVirtualNicDispatcher::start( + Box::new(tunnel), + member_tunnel_table.clone(), + valid.clone(), + ); + let close_notifier = Arc::new(Notify::new()); + let (_member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); + + member_tunnel_table + .register( + member_id, + uuid::Uuid::from_u128(1), + shared_tunnel, + close_notifier.clone(), + ) + .unwrap(); + let empty_claims = SharedIfConfigClaims::default(); + dispatcher + .update_sources(member_id, &empty_claims) + .await + .unwrap(); + + tun_tx.send(Err(TunnelError::Shutdown)).unwrap(); + + tokio::time::timeout(Duration::from_secs(1), close_notifier.notified()) + .await + .unwrap(); + assert!(!valid.load(Ordering::Acquire)); + assert!(member_tunnel_table.dispatcher_channels().is_none()); + } + + #[tokio::test] + async fn dispatcher_shutdown_for_replacement_keeps_shared_nic_valid() { + let member_id = uuid::Uuid::from_u128(1); + let (_tun_tx, tun_rx) = mpsc::unbounded_channel(); + let tun_stream = futures::stream::unfold(tun_rx, |mut receiver| async move { + receiver.recv().await.map(|packet| (packet, receiver)) + }); + let tun_sink = futures::sink::unfold((), |(), _packet: ZCPacket| async { + Ok::<(), TunnelError>(()) + }); + let tunnel = TunnelWrapper::new(tun_stream, tun_sink, None); + let member_tunnel_table = SharedVirtualNicMemberTunnelTable::default(); + let valid = Arc::new(AtomicBool::new(true)); + let dispatcher = SharedVirtualNicDispatcher::start( + Box::new(tunnel), + member_tunnel_table.clone(), + valid.clone(), + ); + let close_notifier = Arc::new(Notify::new()); + let (_member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); + + member_tunnel_table + .register( + member_id, + uuid::Uuid::from_u128(1), + shared_tunnel, + close_notifier.clone(), + ) + .unwrap(); + let empty_claims = SharedIfConfigClaims::default(); + dispatcher + .update_sources(member_id, &empty_claims) + .await + .unwrap(); + + dispatcher.shutdown_without_invalidation().await; + + tokio::time::timeout(Duration::from_secs(1), close_notifier.notified()) + .await + .unwrap(); + assert!(valid.load(Ordering::Acquire)); + assert!(member_tunnel_table.dispatcher_channels().is_none()); + } +} diff --git a/easytier/src/instance/test_instance.rs b/easytier/src/instance/test_instance.rs index 3d815950..be52ba55 100644 --- a/easytier/src/instance/test_instance.rs +++ b/easytier/src/instance/test_instance.rs @@ -7,6 +7,8 @@ use easytier_core::{ process_runtime::CoreProcessRuntime, }; +#[cfg(feature = "tun")] +use crate::instance::shared_virtual_nic::{ArcSharedVirtualNicRegistry, SharedVirtualNicRegistry}; use crate::{ common::global_ctx::{ArcGlobalCtx, GlobalCtx}, instance::{ @@ -15,6 +17,8 @@ use crate::{ }, socket::udp::RuntimeUdpSocket, }; +#[cfg(feature = "tun")] +use tokio::sync::Mutex; pub(crate) struct TestInstance { core: Arc, @@ -26,7 +30,27 @@ impl TestInstance { config: TomlConfig, process_runtime: Arc, ) -> Self { - Self::compose(config, process_runtime, |_| {}) + Self::compose( + config, + process_runtime, + #[cfg(feature = "tun")] + Arc::new(Mutex::new(SharedVirtualNicRegistry::new())), + |_| {}, + ) + } + + #[cfg(feature = "tun")] + pub fn new_with_process_runtime_and_shared_virtual_nic_registry( + config: TomlConfig, + process_runtime: Arc, + shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, + ) -> Self { + Self::compose(config, process_runtime, shared_virtual_nic_registry, |_| {}) + } + + #[cfg(feature = "tun")] + pub fn new_shared_virtual_nic_registry() -> ArcSharedVirtualNicRegistry { + Arc::new(Mutex::new(SharedVirtualNicRegistry::new())) } pub fn new_with_process_runtime_and_stun_provider( @@ -35,14 +59,21 @@ impl TestInstance { provider: Box>, ) -> Self { let provider: Arc> = Arc::from(provider); - Self::compose(config, process_runtime, move |adapters| { - adapters.replace_stun_provider(provider); - }) + Self::compose( + config, + process_runtime, + #[cfg(feature = "tun")] + Arc::new(Mutex::new(SharedVirtualNicRegistry::new())), + move |adapters| { + adapters.replace_stun_provider(provider); + }, + ) } fn compose( config: TomlConfig, process_runtime: Arc, + #[cfg(feature = "tun")] shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, customize: impl FnOnce( &mut easytier_core::instance::CoreHostAdapters< crate::instance::host::NativeInstanceHost, @@ -50,7 +81,11 @@ impl TestInstance { ), ) -> Self { let global_ctx = Arc::new(GlobalCtx::new(config.clone())); - let runtime_host = NativeInstanceRuntimeHost::new(global_ctx.clone()); + let runtime_host = NativeInstanceRuntimeHost::new( + global_ctx.clone(), + #[cfg(feature = "tun")] + shared_virtual_nic_registry, + ); let mut adapters = runtime_core_host_adapters_with_packet_egress( global_ctx.clone(), process_runtime, diff --git a/easytier/src/instance/virtual_nic.rs b/easytier/src/instance/virtual_nic.rs index 43c97aa6..f223c5a6 100644 --- a/easytier/src/instance/virtual_nic.rs +++ b/easytier/src/instance/virtual_nic.rs @@ -11,6 +11,7 @@ use crate::common::{ error::Error, global_ctx::{ArcGlobalCtx, GlobalCtxEvent}, ifcfg::{IfConfiger, IfConfiguerTrait}, + netns::NetNS, }; use easytier_core::{ @@ -43,6 +44,10 @@ use zerocopy::{NativeEndian, NetworkEndian}; #[cfg(target_os = "windows")] use crate::common::ifcfg::RegistryManager; +use super::shared_virtual_nic::{ + ArcSharedVirtualNicRegistry, SharedVirtualNicMember, SharedVirtualNicMemberId, +}; + pin_project! { pub struct TunStream { #[pin] @@ -240,8 +245,32 @@ impl AsyncWrite for TunAsyncWrite { } } +pub struct VirtualNicConfig { + dev_name: String, + mtu: u32, + net_ns: NetNS, +} + +impl VirtualNicConfig { + pub fn new(dev_name: String, mtu: u32, net_ns: NetNS) -> Self { + Self { + dev_name, + mtu, + net_ns, + } + } + + pub fn mtu(&self) -> u32 { + self.mtu + } + + pub fn net_ns_name(&self) -> Option { + self.net_ns.name() + } +} + pub struct VirtualNic { - global_ctx: ArcGlobalCtx, + config: VirtualNicConfig, ifname: Option, ifcfg: Box, @@ -266,11 +295,11 @@ impl Drop for VirtualNic { } impl VirtualNic { - pub fn new(global_ctx: ArcGlobalCtx) -> Self { + pub fn new(config: VirtualNicConfig) -> Self { Self { - global_ctx, + config, ifname: None, - ifcfg: Box::new(IfConfiger {}), + ifcfg: Box::new(IfConfiger::default()), } } @@ -487,14 +516,14 @@ impl VirtualNic { Ok(()) } - async fn create_tun(&self) -> Result { + async fn create_tun(&mut self) -> Result { let mut config = Configuration::default(); config.layer(Layer::L3); // FreeBSD specific: Check and restore TUN interfaces before creating new one #[cfg(target_os = "freebsd")] { - let dev_name = self.global_ctx.get_flags().dev_name; + let dev_name = self.config.dev_name.clone(); if !dev_name.is_empty() { // Restore TUN interface name if needed, ignoring errors as it's not critical @@ -507,7 +536,7 @@ impl VirtualNic { // Check and create TUN device node if necessary (Linux only) Self::ensure_tun_device_node().await; - let dev_name = self.global_ctx.get_flags().dev_name; + let dev_name = self.config.dev_name.clone(); if !dev_name.is_empty() { config.tun_name(&dev_name); } @@ -521,7 +550,7 @@ impl VirtualNic { #[cfg(target_os = "windows")] { - let dev_name = self.global_ctx.get_flags().dev_name; + let dev_name = self.config.dev_name.clone(); match crate::arch::windows::add_self_to_firewall_allowlist() { Ok(_) => tracing::info!("add_self_to_firewall_allowlist successful!"), @@ -553,10 +582,7 @@ impl VirtualNic { let random_dev_name = format!("et_{}_{}", c, s); config.tun_name(random_dev_name.clone()); - - let mut flags = self.global_ctx.get_flags(); - flags.dev_name = random_dev_name.clone(); - self.global_ctx.set_flags(flags); + self.config.dev_name = random_dev_name; } config.platform_config(|config| { @@ -570,10 +596,15 @@ impl VirtualNic { config.up(); - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); Ok(tun::create(&config)?) } + #[cfg(mobile)] + pub fn set_mobile_tun_fd_name(&mut self, tun_fd: std::os::fd::RawFd) { + self.ifname = Some(format!("tunfd_{}", tun_fd)); + } + #[cfg(mobile)] pub async fn create_dev_for_mobile( &mut self, @@ -609,7 +640,7 @@ impl VirtualNic { None, ); - self.ifname = Some(format!("tunfd_{}", tun_fd)); + self.set_mobile_tun_fd_name(tun_fd); Ok(Box::new(ft)) } @@ -627,7 +658,7 @@ impl VirtualNic { // FreeBSD TUN interface rename functionality #[cfg(target_os = "freebsd")] { - let dev_name = self.global_ctx.get_flags().dev_name; + let dev_name = self.config.dev_name.clone(); if !dev_name.is_empty() && dev_name != ifname { // Use ifconfig to rename the TUN interface @@ -663,15 +694,10 @@ impl VirtualNic { let dev = AsyncDevice::new(dev)?; - let flags = self.global_ctx.get_flags(); - let mut mtu_in_config = flags.mtu; - if flags.enable_encryption { - mtu_in_config -= 20; - } { // set mtu by ourselves, rust-tun does not handle it correctly on windows - let _g = self.global_ctx.net_ns.guard(); - self.ifcfg.set_mtu(ifname.as_str(), mtu_in_config).await?; + let _g = self.config.net_ns.guard(); + self.ifcfg.set_mtu(ifname.as_str(), self.config.mtu).await?; } let has_packet_info = cfg!(all(target_os = "macos", not(feature = "macos-ne"))); @@ -718,15 +744,63 @@ impl VirtualNic { } pub async fn link_up(&self) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg.set_link_status(self.ifname(), true).await?; Ok(()) } pub async fn add_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + self.add_route_with_cost(address, cidr, None).await + } + + pub async fn add_route_with_cost( + &self, + address: Ipv4Addr, + cidr: u8, + cost: Option, + ) -> Result<(), Error> { + self.add_route_with_cost_and_source_hint(address, cidr, cost, None) + .await + } + + pub async fn add_route_with_cost_and_source_hint( + &self, + address: Ipv4Addr, + cidr: u8, + cost: Option, + source_hint: Option, + ) -> Result<(), Error> { + let _g = self.config.net_ns.guard(); self.ifcfg - .add_ipv4_route(self.ifname(), address, cidr, None) + .add_ipv4_route_with_source_hint(self.ifname(), address, cidr, cost, source_hint) + .await?; + Ok(()) + } + + pub async fn remove_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { + let _g = self.config.net_ns.guard(); + self.ifcfg + .remove_ipv4_route(self.ifname(), address, cidr) + .await?; + Ok(()) + } + + pub async fn remove_route_with_cost_and_source_hint( + &self, + address: Ipv4Addr, + cidr: u8, + cost: Option, + source_hint: Option, + ) -> Result<(), Error> { + let _g = self.config.net_ns.guard(); + self.ifcfg + .remove_ipv4_route_with_cost_and_source_hint( + self.ifname(), + address, + cidr, + cost, + source_hint, + ) .await?; Ok(()) } @@ -741,7 +815,7 @@ impl VirtualNic { cidr: u8, cost: Option, ) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg .add_ipv6_route(self.ifname(), address, cidr, cost) .await?; @@ -749,7 +823,7 @@ impl VirtualNic { } pub async fn remove_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg .remove_ipv6_route(self.ifname(), address, cidr) .await?; @@ -757,19 +831,19 @@ impl VirtualNic { } pub async fn remove_ip(&self, ip: Option) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg.remove_ip(self.ifname(), ip).await?; Ok(()) } pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg.remove_ipv6(self.ifname(), ip).await?; Ok(()) } pub async fn add_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg .add_ipv4_ip(self.ifname(), ip, cidr as u8) .await?; @@ -777,15 +851,208 @@ impl VirtualNic { } pub async fn add_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { - let _g = self.global_ctx.net_ns.guard(); + let _g = self.config.net_ns.guard(); self.ifcfg .add_ipv6_ip(self.ifname(), ip, cidr as u8) .await?; Ok(()) } - pub fn get_ifcfg(&self) -> impl IfConfiguerTrait + use<> { - IfConfiger {} + pub async fn set_mtu(&self, mtu: u32) -> Result<(), Error> { + let _g = self.config.net_ns.guard(); + self.ifcfg.set_mtu(self.ifname(), mtu).await?; + Ok(()) + } + + pub fn configured_mtu(&self) -> u32 { + self.config.mtu + } + + #[cfg(test)] + pub(crate) fn set_ifname_for_test(&mut self, ifname: String) { + self.ifname = Some(ifname); + } + + #[cfg(test)] + pub(crate) fn set_ifcfg_for_test( + &mut self, + ifcfg: Box, + ) { + self.ifcfg = ifcfg; + } + + pub fn get_ifcfg(&self) -> IfConfiger { + IfConfiger::default() + } +} + +#[derive(Clone)] +pub enum NicBackend { + Dedicated(Arc>), + Shared(SharedVirtualNicMember), +} + +impl NicBackend { + pub fn dedicated(nic: Arc>) -> Self { + Self::Dedicated(nic) + } + + pub fn shared(member: SharedVirtualNicMember) -> Self { + Self::Shared(member) + } + + pub async fn create_dev(&self) -> Result, Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.create_dev().await, + Self::Shared(member) => member.create_dev().await, + } + } + + #[cfg(mobile)] + pub async fn create_dev_for_mobile( + &self, + tun_fd: std::os::fd::RawFd, + replace_tun_fd: bool, + ) -> Result, Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.create_dev_for_mobile(tun_fd).await, + Self::Shared(member) => member.create_dev_for_mobile(tun_fd, replace_tun_fd).await, + } + } + + pub async fn ifname(&self) -> Option { + match self { + Self::Dedicated(nic) => nic + .lock() + .await + .ifname + .as_ref() + .map(|ifname| ifname.to_owned()), + Self::Shared(member) => { + let shared_nic = member.shared_nic(); + let nic = { + let shared_nic = shared_nic.lock().await; + shared_nic.nic() + }; + nic.lock() + .await + .ifname + .as_ref() + .map(|ifname| ifname.to_owned()) + } + } + } + + #[cfg(not(target_os = "linux"))] + /// Returns a raw ifcfg handle and interface name for platform cleanup. + /// + /// This does not carry `VirtualNic`'s netns guard. Use the typed + /// `NicBackend` methods for normal IP and route configuration. + pub async fn ifcfg_and_ifname(&self) -> Result<(IfConfiger, String), Error> { + match self { + Self::Dedicated(nic) => { + let nic = nic.lock().await; + Ok((nic.get_ifcfg(), nic.ifname().to_owned())) + } + Self::Shared(member) => member.ifcfg_and_ifname().await, + } + } + + pub async fn link_up(&self) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.link_up().await, + Self::Shared(member) => member.link_up().await, + } + } + + pub async fn add_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.add_route(address, cidr).await, + Self::Shared(member) => member.add_route(address, cidr).await, + } + } + + pub async fn add_route_with_cost( + &self, + address: Ipv4Addr, + cidr: u8, + cost: Option, + ) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => { + nic.lock() + .await + .add_route_with_cost(address, cidr, cost) + .await + } + Self::Shared(member) => member.add_route_with_cost(address, cidr, cost).await, + } + } + + pub async fn remove_route(&self, address: Ipv4Addr, cidr: u8) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.remove_route(address, cidr).await, + Self::Shared(member) => member.remove_route(address, cidr).await, + } + } + + pub async fn add_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.add_ipv6_route(address, cidr).await, + Self::Shared(member) => member.add_ipv6_route(address, cidr).await, + } + } + + pub async fn add_ipv6_route_with_cost( + &self, + address: Ipv6Addr, + cidr: u8, + cost: Option, + ) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => { + nic.lock() + .await + .add_ipv6_route_with_cost(address, cidr, cost) + .await + } + Self::Shared(member) => member.add_ipv6_route_with_cost(address, cidr, cost).await, + } + } + + pub async fn remove_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.remove_ipv6_route(address, cidr).await, + Self::Shared(member) => member.remove_ipv6_route(address, cidr).await, + } + } + + pub async fn remove_ip(&self, ip: Option) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.remove_ip(ip).await, + Self::Shared(member) => member.remove_ip(ip).await, + } + } + + pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.remove_ipv6(ip).await, + Self::Shared(member) => member.remove_ipv6(ip).await, + } + } + + pub async fn add_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.add_ip(ip, cidr).await, + Self::Shared(member) => member.add_ip(ip, cidr).await, + } + } + + pub async fn add_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { + match self { + Self::Dedicated(nic) => nic.lock().await.add_ipv6(ip, cidr).await, + Self::Shared(member) => member.add_ipv6(ip, cidr).await, + } } } @@ -796,7 +1063,7 @@ pub struct NicCtx { close_notifier: Arc, - nic: Arc>, + backend: NicBackend, tasks: JoinSet<()>, #[cfg(target_os = "windows")] @@ -804,11 +1071,47 @@ pub struct NicCtx { } impl NicCtx { - pub(crate) fn new( + fn virtual_nic_config_from_parts( + dev_name: String, + mut mtu: u32, + enable_encryption: bool, + net_ns: NetNS, + ) -> VirtualNicConfig { + if enable_encryption { + mtu -= 20; + } + + VirtualNicConfig::new(dev_name, mtu, net_ns) + } + + fn virtual_nic_config(global_ctx: &ArcGlobalCtx) -> VirtualNicConfig { + let flags = global_ctx.get_flags(); + Self::virtual_nic_config_from_parts( + flags.dev_name, + flags.mtu, + flags.enable_encryption, + global_ctx.net_ns.clone(), + ) + } + + pub(crate) fn shared_route_backend_for_dns(&self) -> Option { + match self.backend { + NicBackend::Dedicated(_) => None, + NicBackend::Shared(_) => Some(self.backend.clone()), + } + } + + fn dedicated_backend(global_ctx: &ArcGlobalCtx) -> NicBackend { + let nic_config = Self::virtual_nic_config(global_ctx); + NicBackend::dedicated(Arc::new(Mutex::new(VirtualNic::new(nic_config)))) + } + + fn new_with_backend( global_ctx: ArcGlobalCtx, packet_plane: Arc, peer_packet_receiver: Arc>, close_notifier: Arc, + backend: NicBackend, ) -> Self { NicCtx { global_ctx: global_ctx.clone(), @@ -817,7 +1120,7 @@ impl NicCtx { close_notifier, - nic: Arc::new(Mutex::new(VirtualNic::new(global_ctx))), + backend, tasks: JoinSet::new(), #[cfg(target_os = "windows")] @@ -825,40 +1128,110 @@ impl NicCtx { } } + pub(crate) fn new( + global_ctx: ArcGlobalCtx, + packet_plane: Arc, + peer_packet_receiver: Arc>, + close_notifier: Arc, + ) -> Self { + let backend = Self::dedicated_backend(&global_ctx); + + Self::new_with_backend( + global_ctx, + packet_plane, + peer_packet_receiver, + close_notifier, + backend, + ) + } + + pub(crate) async fn new_shared( + global_ctx: ArcGlobalCtx, + packet_plane: Arc, + peer_packet_receiver: Arc>, + close_notifier: Arc, + registry: ArcSharedVirtualNicRegistry, + member_id: SharedVirtualNicMemberId, + ) -> Result { + let flags = global_ctx.get_flags(); + #[cfg(mobile)] + let dev_name = String::new(); + #[cfg(not(mobile))] + let dev_name = flags.dev_name.clone(); + #[cfg(not(mobile))] + if dev_name.is_empty() { + return Err(anyhow::anyhow!("shared virtual nic requires dev_name").into()); + } + #[cfg(mobile)] + let net_ns = NetNS::new(None); + #[cfg(not(mobile))] + let net_ns = global_ctx.net_ns.clone(); + let nic_config = Self::virtual_nic_config_from_parts( + dev_name.clone(), + flags.mtu, + flags.enable_encryption, + net_ns, + ); + + let member = registry.lock().await.create_member( + dev_name, + nic_config, + member_id, + close_notifier.clone(), + ); + let backend = NicBackend::shared(member); + + Ok(Self::new_with_backend( + global_ctx, + packet_plane, + peer_packet_receiver, + close_notifier, + backend, + )) + } + pub async fn ifname(&self) -> Option { - let nic = self.nic.lock().await; - nic.ifname.as_ref().map(|s| s.to_owned()) + self.backend.ifname().await + } + + async fn tun_ifname(&self) -> Result { + self.backend + .ifname() + .await + .ok_or_else(|| anyhow::anyhow!("tun device has no interface name").into()) } pub async fn assign_ipv4_to_tun_device(&self, ipv4_addr: cidr::Ipv4Inet) -> Result<(), Error> { - let nic = self.nic.lock().await; - nic.link_up().await?; - nic.remove_ip(None).await?; - nic.add_ip(ipv4_addr.address(), ipv4_addr.network_length() as i32) + self.backend.link_up().await?; + self.backend.remove_ip(None).await?; + self.backend + .add_ip(ipv4_addr.address(), ipv4_addr.network_length() as i32) .await?; #[cfg(any( all(target_os = "macos", not(feature = "macos-ne")), target_os = "freebsd" ))] { - nic.add_route(ipv4_addr.first_address(), ipv4_addr.network_length()) + self.backend + .add_route(ipv4_addr.first_address(), ipv4_addr.network_length()) .await?; } Ok(()) } pub async fn assign_ipv6_to_tun_device(&self, ipv6_addr: cidr::Ipv6Inet) -> Result<(), Error> { - let nic = self.nic.lock().await; - nic.link_up().await?; - nic.remove_ipv6(None).await?; - nic.add_ipv6(ipv6_addr.address(), ipv6_addr.network_length() as i32) + self.backend.link_up().await?; + self.backend.remove_ipv6(None).await?; + self.backend + .add_ipv6(ipv6_addr.address(), ipv6_addr.network_length() as i32) .await?; #[cfg(any( all(target_os = "macos", not(feature = "macos-ne")), target_os = "freebsd" ))] { - nic.add_ipv6_route(ipv6_addr.first_address(), ipv6_addr.network_length()) + self.backend + .add_ipv6_route(ipv6_addr.first_address(), ipv6_addr.network_length()) .await?; } Ok(()) @@ -923,6 +1296,15 @@ impl NicCtx { }); } + fn start_tunnel_forwarding(&mut self, tunnel: Box) -> Result<(), Error> { + let (stream, sink) = tunnel.split(); + + self.do_forward_nic_to_peers_task(stream)?; + self.do_forward_peers_to_nic(sink); + + Ok(()) + } + #[cfg(target_os = "windows")] fn start_windows_udp_broadcast_relay(&mut self, virtual_ipv4: Ipv4Inet) { if !self.global_ctx.get_flags().enable_udp_broadcast_relay { @@ -948,13 +1330,11 @@ impl NicCtx { } async fn apply_route_changes( - ifcfg: &impl IfConfiguerTrait, - ifname: &str, - net_ns: &crate::common::netns::NetNS, + backend: &NicBackend, cur_proxy_cidrs: &mut BTreeSet, added: Vec, removed: Vec, - ) { + ) -> Result<(), Error> { tracing::debug!(?added, ?removed, "applying proxy_cidrs route changes"); // Remove routes @@ -962,9 +1342,8 @@ impl NicCtx { if !cur_proxy_cidrs.contains(&cidr) { continue; } - let _g = net_ns.guard(); - let ret = ifcfg - .remove_ipv4_route(ifname, cidr.first_address(), cidr.network_length()) + let ret = backend + .remove_route(cidr.first_address(), cidr.network_length()) .await; if ret.is_err() { @@ -978,30 +1357,31 @@ impl NicCtx { } // Add routes + let mut first_error = None; for cidr in added { if cur_proxy_cidrs.contains(&cidr) { continue; } - let _g = net_ns.guard(); - let ret = ifcfg - .add_ipv4_route(ifname, cidr.first_address(), cidr.network_length(), None) - .await; - - if ret.is_err() { - tracing::trace!( - cidr = ?cidr, - err = ?ret, - "add route failed.", - ); + match backend + .add_route(cidr.first_address(), cidr.network_length()) + .await + { + Ok(()) => { + cur_proxy_cidrs.insert(cidr); + } + Err(err) => { + tracing::error!(?cidr, ?err, "add route failed"); + if first_error.is_none() { + first_error = Some(err); + } + } } - cur_proxy_cidrs.insert(cidr); } + first_error.map_or(Ok(()), Err) } async fn apply_public_ipv6_route_changes( - ifcfg: &impl IfConfiguerTrait, - ifname: &str, - net_ns: &crate::common::netns::NetNS, + backend: &NicBackend, cur_routes: &mut BTreeSet, added: Vec, removed: Vec, @@ -1010,9 +1390,8 @@ impl NicCtx { if !cur_routes.contains(&route) { continue; } - let _g = net_ns.guard(); - let ret = ifcfg - .remove_ipv6_route(ifname, route.address(), route.network_length()) + let ret = backend + .remove_ipv6_route(route.address(), route.network_length()) .await; if ret.is_err() { tracing::trace!(route = ?route, err = ?ret, "remove public ipv6 route failed"); @@ -1024,9 +1403,8 @@ impl NicCtx { if cur_routes.contains(&route) { continue; } - let _g = net_ns.guard(); - let ret = ifcfg - .add_ipv6_route(ifname, route.address(), route.network_length(), None) + let ret = backend + .add_ipv6_route(route.address(), route.network_length()) .await; if ret.is_err() { tracing::trace!(route = ?route, err = ?ret, "add public ipv6 route failed"); @@ -1039,33 +1417,30 @@ impl NicCtx { async fn run_proxy_cidrs_route_updater(&mut self) -> Result<(), Error> { let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); - let net_ns = self.global_ctx.net_ns.clone(); - let nic = self.nic.lock().await; - let ifcfg = nic.get_ifcfg(); - let ifname = nic.ifname().to_owned(); + let backend = self.backend.clone(); let mut event_receiver = global_ctx.subscribe(); + let mut cur_proxy_cidrs = BTreeSet::::new(); + + // Initial sync: get current proxy_cidrs state and apply routes + let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await else { + tracing::error!("proxy CIDR monitor host is unavailable"); + return Ok(()); + }; + if let Err(err) = + Self::apply_route_changes(&backend, &mut cur_proxy_cidrs, diff.added, diff.removed) + .await + { + if matches!(&backend, NicBackend::Shared(_)) { + return Err(err); + } + } + self.tasks.spawn(async move { - let mut cur_proxy_cidrs = BTreeSet::::new(); - - // Initial sync: get current proxy_cidrs state and apply routes - let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await else { - tracing::error!("proxy CIDR monitor host is unavailable"); - return; - }; - Self::apply_route_changes( - &ifcfg, - &ifname, - &net_ns, - &mut cur_proxy_cidrs, - diff.added, - diff.removed, - ) - .await; - loop { - let event = match event_receiver.recv().await { - Ok(event) => event, + let should_sync = match event_receiver.recv().await { + Ok(GlobalCtxEvent::ProxyCidrsUpdated(_, _)) => true, + Ok(_) => false, Err(tokio::sync::broadcast::error::RecvError::Closed) => { tracing::debug!("event bus closed, stopping proxy_cidrs route updater"); break; @@ -1075,31 +1450,29 @@ impl NicCtx { "event bus lagged in proxy_cidrs route updater, doing full sync" ); event_receiver = event_receiver.resubscribe(); - // Full sync after lagged to recover consistent state - let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await - else { - tracing::error!("proxy CIDR monitor host is unavailable"); - return; - }; - GlobalCtxEvent::ProxyCidrsUpdated(diff.added, diff.removed) + true } }; + if !should_sync { + continue; + } - // Only handle ProxyCidrsUpdated events - let (added, removed) = match event { - GlobalCtxEvent::ProxyCidrsUpdated(added, removed) => (added, removed), - _ => continue, + // Full sync also retries routes that previously failed to apply. + let Some(diff) = packet_plane.proxy_cidr_diff(&cur_proxy_cidrs).await else { + tracing::error!("proxy CIDR monitor host is unavailable"); + return; }; - Self::apply_route_changes( - &ifcfg, - &ifname, - &net_ns, + if let Err(err) = Self::apply_route_changes( + &backend, &mut cur_proxy_cidrs, - added, - removed, + diff.added, + diff.removed, ) - .await; + .await + { + tracing::error!(?err, "failed to update proxy CIDR routes"); + } } }); @@ -1109,10 +1482,7 @@ impl NicCtx { async fn run_public_ipv6_route_updater(&mut self) -> Result<(), Error> { let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); - let net_ns = self.global_ctx.net_ns.clone(); - let nic = self.nic.lock().await; - let ifcfg = nic.get_ifcfg(); - let ifname = nic.ifname().to_owned(); + let backend = self.backend.clone(); let mut event_receiver = global_ctx.subscribe(); self.tasks.spawn(async move { @@ -1120,9 +1490,7 @@ impl NicCtx { let initial_routes = packet_plane.public_ipv6_routes().await; let initial_added = initial_routes.iter().copied().collect::>(); Self::apply_public_ipv6_route_changes( - &ifcfg, - &ifname, - &net_ns, + &backend, &mut cur_routes, initial_added, Vec::new(), @@ -1147,15 +1515,8 @@ impl NicCtx { _ => continue, }; - Self::apply_public_ipv6_route_changes( - &ifcfg, - &ifname, - &net_ns, - &mut cur_routes, - added, - removed, - ) - .await; + Self::apply_public_ipv6_route_changes(&backend, &mut cur_routes, added, removed) + .await; } }); @@ -1165,20 +1526,22 @@ impl NicCtx { async fn run_public_ipv6_addr_updater(&mut self) -> Result<(), Error> { let packet_plane = self.packet_plane.clone(); let global_ctx = self.global_ctx.clone(); - let nic = self.nic.clone(); + let backend = self.backend.clone(); let mut event_receiver = global_ctx.subscribe(); self.tasks.spawn(async move { let mut current_addr = packet_plane.public_ipv6_addr().await; if let Some(addr) = current_addr { - let nic = nic.lock().await; - if let Err(err) = nic.link_up().await { + if let Err(err) = backend.link_up().await { tracing::warn!(?err, "failed to bring public ipv6 nic link up"); } - if let Err(err) = nic.add_ipv6(addr.address(), addr.network_length() as i32).await { + if let Err(err) = backend + .add_ipv6(addr.address(), addr.network_length() as i32) + .await + { tracing::warn!(addr = ?addr, ?err, "failed to add public ipv6 address"); } - if let Err(err) = nic + if let Err(err) = backend .add_ipv6_route_with_cost(Ipv6Addr::UNSPECIFIED, 0, Some(5)) .await { @@ -1203,24 +1566,28 @@ impl NicCtx { }; current_addr = new; - let nic = nic.lock().await; - if let Err(err) = nic.link_up().await { + if let Err(err) = backend.link_up().await { tracing::warn!(?err, "failed to bring public ipv6 nic link up"); } if let Some(old) = old { - if let Err(err) = nic.remove_ipv6_route(Ipv6Addr::UNSPECIFIED, 0).await { + if let Err(err) = backend + .remove_ipv6_route(Ipv6Addr::UNSPECIFIED, 0) + .await + { tracing::warn!(route = %Ipv6Addr::UNSPECIFIED, prefix = 0, ?err, "failed to remove default public ipv6 route"); } - if let Err(err) = nic.remove_ipv6(Some(old)).await { + if let Err(err) = backend.remove_ipv6(Some(old)).await { tracing::warn!(addr = ?old, ?err, "failed to remove old public ipv6 address"); } } if let Some(new) = new { - if let Err(err) = nic.add_ipv6(new.address(), new.network_length() as i32).await + if let Err(err) = backend + .add_ipv6(new.address(), new.network_length() as i32) + .await { tracing::warn!(addr = ?new, ?err, "failed to add public ipv6 address"); } - if let Err(err) = nic + if let Err(err) = backend .add_ipv6_route_with_cost(Ipv6Addr::UNSPECIFIED, 0, Some(5)) .await { @@ -1238,43 +1605,42 @@ impl NicCtx { ipv4_addr: Option, ipv6_addr: Option, ) -> Result<(), Error> { - let tunnel = { - let mut nic = self.nic.lock().await; - match nic.create_dev().await { - Ok(ret) => { - #[cfg(target_os = "windows")] - { - let dev_name = self.global_ctx.get_flags().dev_name; - let _ = RegistryManager::reg_change_catrgory_in_profile(&dev_name); - } + let tunnel = match self.backend.create_dev().await { + Ok(ret) => { + let ifname = self.tun_ifname().await?; - #[cfg(any( - all(target_os = "macos", not(feature = "macos-ne")), - target_os = "freebsd" - ))] - { - // remove the 10.0.0.0/24 route (which is added by rust-tun by default) - let _ = nic - .ifcfg - .remove_ipv4_route(nic.ifname(), "10.0.0.0".parse().unwrap(), 24) - .await; + #[cfg(target_os = "windows")] + { + let mut flags = self.global_ctx.get_flags(); + if flags.dev_name.is_empty() { + flags.dev_name = ifname.clone(); + self.global_ctx.set_flags(flags); } + let _ = RegistryManager::reg_change_catrgory_in_profile(&ifname); + } - self.global_ctx - .set_tun_device_ready(nic.ifname().to_string()); - ret - } - Err(err) => { - self.global_ctx.set_tun_device_error(err.to_string()); - return Err(err); + #[cfg(any( + all(target_os = "macos", not(feature = "macos-ne")), + target_os = "freebsd" + ))] + { + // remove the 10.0.0.0/24 route (which is added by rust-tun by default) + let (ifcfg, ifname) = self.backend.ifcfg_and_ifname().await?; + let _ = ifcfg + .remove_ipv4_route(&ifname, "10.0.0.0".parse().unwrap(), 24) + .await; } + + self.global_ctx.set_tun_device_ready(ifname); + ret + } + Err(err) => { + self.global_ctx.set_tun_device_error(err.to_string()); + return Err(err); } }; - let (stream, sink) = tunnel.split(); - - self.do_forward_nic_to_peers_task(stream)?; - self.do_forward_peers_to_nic(sink); + self.start_tunnel_forwarding(tunnel)?; // Assign IPv4 address if provided if let Some(ipv4_addr) = ipv4_addr { @@ -1298,26 +1664,34 @@ impl NicCtx { } #[cfg(mobile)] - pub async fn run_for_mobile(&mut self, tun_fd: std::os::fd::RawFd) -> Result<(), Error> { - let tunnel = { - let mut nic = self.nic.lock().await; - match nic.create_dev_for_mobile(tun_fd).await { - Ok(ret) => { - self.global_ctx - .set_tun_device_ready(nic.ifname().to_string()); - ret - } - Err(err) => { - self.global_ctx.set_tun_device_error(err.to_string()); - return Err(err); - } + pub async fn run_for_mobile( + &mut self, + tun_fd: std::os::fd::RawFd, + replace_tun_fd: bool, + ) -> Result<(), Error> { + let (tunnel, ifname) = match self + .backend + .create_dev_for_mobile(tun_fd, replace_tun_fd) + .await + { + Ok(ret) => { + let ifname = self.tun_ifname().await?; + (ret, ifname) + } + Err(err) => { + self.global_ctx.set_tun_device_error(err.to_string()); + return Err(err); } }; - let (stream, sink) = tunnel.split(); + if let Some(ipv4_addr) = self.global_ctx.get_ipv4() { + self.assign_ipv4_to_tun_device(ipv4_addr).await?; + } + self.run_proxy_cidrs_route_updater().await?; - self.do_forward_nic_to_peers_task(stream)?; - self.do_forward_peers_to_nic(sink); + self.global_ctx.set_tun_device_ready(ifname); + + self.start_tunnel_forwarding(tunnel)?; Ok(()) } @@ -1327,10 +1701,11 @@ impl NicCtx { mod tests { use crate::common::{error::Error, global_ctx::tests::get_mock_global_ctx}; - use super::VirtualNic; + use super::{NicCtx, VirtualNic}; async fn run_test_helper() -> Result { - let mut dev = VirtualNic::new(get_mock_global_ctx()); + let global_ctx = get_mock_global_ctx(); + let mut dev = VirtualNic::new(NicCtx::virtual_nic_config(&global_ctx)); let _tunnel = dev.create_dev().await?; tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; diff --git a/easytier/src/tests/mod.rs b/easytier/src/tests/mod.rs index 623ca559..728c9705 100644 --- a/easytier/src/tests/mod.rs +++ b/easytier/src/tests/mod.rs @@ -1,6 +1,9 @@ #[cfg(target_os = "linux")] mod three_node; +#[cfg(all(target_os = "linux", feature = "tun"))] +mod shared_virtual_nic; + mod ipv6_test; #[cfg(target_os = "linux")] diff --git a/easytier/src/tests/shared_virtual_nic.rs b/easytier/src/tests/shared_virtual_nic.rs new file mode 100644 index 00000000..f47e52a2 --- /dev/null +++ b/easytier/src/tests/shared_virtual_nic.rs @@ -0,0 +1,476 @@ +use std::{net::Ipv4Addr, process::Command, sync::Arc, time::Duration}; + +use easytier_core::{config::PeerId, process_runtime::CoreProcessRuntime}; + +use super::{ + InstanceTestExt as _, add_ns_to_bridge, create_netns, del_netns, drop_insts, ping_test, + prepare_bridge, +}; +use crate::{ + common::{ + config::{ConfigLoader, NetworkIdentity, TomlConfigLoader}, + netns::{NetNS, ROOT_NETNS_NAME}, + }, + instance::{ + shared_virtual_nic::{ArcSharedVirtualNicRegistry, SharedIpv4Route}, + test_instance::TestInstance as Instance, + }, + tunnel::common::tests::wait_for_condition, +}; + +const PROXY_CIDR: &str = "10.1.2.0/24"; +const WAIT: Duration = Duration::from_secs(10); + +#[derive(Clone)] +struct SharedTestRuntime { + process: Arc, + registry: ArcSharedVirtualNicRegistry, +} + +impl SharedTestRuntime { + fn new() -> Self { + Self { + process: CoreProcessRuntime::new(), + registry: Instance::new_shared_virtual_nic_registry(), + } + } + + fn instance(&self, config: TomlConfigLoader) -> Instance { + Instance::new_with_process_runtime_and_shared_virtual_nic_registry( + config, + self.process.clone(), + self.registry.clone(), + ) + } +} + +fn test_config( + instance_name: &str, + network_name: &str, + network_secret: &str, + netns: Option<&str>, + dev_name: Option<&str>, + ipv4: &str, +) -> TomlConfigLoader { + let config = TomlConfigLoader::default(); + config.set_inst_name(instance_name.to_owned()); + config.set_network_identity(NetworkIdentity::new( + network_name.to_owned(), + network_secret.to_owned(), + )); + config.set_netns(netns.map(str::to_owned)); + config.set_ipv4(Some(ipv4.parse().unwrap())); + config.set_ipv6(None); + config.set_dhcp(false); + config.set_listeners(vec![]); + config.set_socks5_portal(None); + + let mut flags = config.get_flags(); + flags.dev_name = dev_name.unwrap_or_default().to_owned(); + flags.enable_ipv6 = false; + config.set_flags(flags); + config +} + +fn test_dev_name() -> String { + format!("st{:08x}", rand::random::()) +} + +fn short_name(prefix: &str) -> String { + format!("{prefix}{:04x}", rand::random::()) +} + +struct TestNetnsGuard { + name: String, +} + +impl TestNetnsGuard { + fn new(name: String, ipv4: &str, ipv6: &str) -> Self { + let guard = Self { name }; + del_netns(&guard.name); + create_netns(&guard.name, ipv4, ipv6); + guard + } +} + +impl Drop for TestNetnsGuard { + fn drop(&mut self) { + del_netns(&self.name); + } +} + +struct ProxyLab { + source_ns: String, + owner_ns: String, + target_ns: String, + bridge: String, +} + +impl ProxyLab { + fn new() -> Self { + let suffix = format!("{:04x}", rand::random::()); + let lab = Self { + source_ns: format!("svs{suffix}"), + owner_ns: format!("svo{suffix}"), + target_ns: format!("svt{suffix}"), + bridge: format!("svb{suffix}"), + }; + lab.cleanup(); + + create_netns(&lab.source_ns, "10.1.1.1/24", "fd11::1/64"); + create_netns(&lab.owner_ns, "10.1.2.3/24", "fd12::3/64"); + create_netns(&lab.target_ns, "10.1.2.4/24", "fd12::4/64"); + prepare_bridge(&lab.bridge); + add_ns_to_bridge(&lab.bridge, &lab.owner_ns); + add_ns_to_bridge(&lab.bridge, &lab.target_ns); + lab + } + + fn cleanup(&self) { + del_netns(&self.source_ns); + del_netns(&self.owner_ns); + del_netns(&self.target_ns); + let _ = Command::new("ip") + .args(["link", "del", &self.bridge]) + .output(); + } +} + +impl Drop for ProxyLab { + fn drop(&mut self) { + self.cleanup(); + } +} + +async fn wait_tun_ready(instance: &Instance, expected: &str) { + wait_for_condition( + || async { instance.get_global_ctx().get_tun_device_name().as_deref() == Some(expected) }, + WAIT, + ) + .await; +} + +fn proxy_route_exists( + routes: &[easytier_proto::core_peer::peer::Route], + peer_id: PeerId, + proxy_cidr: &str, +) -> bool { + routes + .iter() + .any(|route| route.peer_id == peer_id && route.proxy_cidrs.iter().any(|c| c == proxy_cidr)) +} + +async fn wait_proxy_route(instance: &Instance, peer_id: PeerId, proxy_cidr: &str) { + wait_for_condition( + || async { + proxy_route_exists( + &instance.get_core_instance().route_snapshots().await, + peer_id, + proxy_cidr, + ) + }, + WAIT, + ) + .await; +} + +async fn wait_proxy_route_absent(instance: &Instance, peer_id: PeerId, proxy_cidr: &str) { + wait_for_condition( + || async { + !proxy_route_exists( + &instance.get_core_instance().route_snapshots().await, + peer_id, + proxy_cidr, + ) + }, + WAIT, + ) + .await; +} + +fn ipv4_route_exists_in_ns(ns: &str, needle: &str) -> bool { + let _root = NetNS::new(Some(ROOT_NETNS_NAME.to_owned())).guard(); + let output = Command::new("ip") + .args(["netns", "exec", ns, "ip", "route", "show"]) + .output() + .unwrap(); + assert!( + output.status.success(), + "failed to list IPv4 routes in {ns}: {}", + String::from_utf8_lossy(&output.stderr) + ); + String::from_utf8_lossy(&output.stdout) + .lines() + .any(|line| line.contains(needle)) +} + +#[cfg(feature = "proxy-cidr-monitor")] +async fn patch_proxy_cidr( + instance: &Instance, + action: crate::proto::api::config::ConfigPatchAction, +) { + use crate::proto::api::config::{InstanceConfigPatch, ProxyNetworkPatch}; + + instance + .get_config_patcher() + .apply_patch(InstanceConfigPatch { + proxy_networks: vec![ProxyNetworkPatch { + action: action as i32, + cidr: Some(PROXY_CIDR.parse().unwrap()), + mapped_cidr: None, + }], + ..Default::default() + }) + .await + .unwrap(); +} + +async fn shared_route_owner_count( + registry: &ArcSharedVirtualNicRegistry, + dev_name: &str, + route: &SharedIpv4Route, +) -> usize { + let nic = { + let registry = registry.lock().await; + registry.get_by_dev_name_for_test(dev_name) + }; + let Some(nic) = nic else { + return 0; + }; + nic.lock().await.ifcfg().owners_of_ipv4_route(route).len() +} + +#[tokio::test] +#[serial_test::serial] +async fn same_namespace_members_share_tun_across_independent_networks() { + let dev_name = test_dev_name(); + let first_peer_ns = TestNetnsGuard::new(short_name("sva"), "10.231.1.2/24", "fd31::2/64"); + let second_peer_ns = TestNetnsGuard::new(short_name("svb"), "10.231.2.2/24", "fd32::2/64"); + let runtime = SharedTestRuntime::new(); + + let mut first = runtime.instance(test_config( + "shared_tun_first", + "shared_tun_network_a", + "shared_tun_secret_a", + None, + Some(&dev_name), + "10.144.250.1/24", + )); + let mut second = runtime.instance(test_config( + "shared_tun_second", + "shared_tun_network_b", + "shared_tun_secret_b", + None, + Some(&dev_name), + "10.144.251.1/24", + )); + let mut first_peer = runtime.instance(test_config( + "shared_tun_first_peer", + "shared_tun_network_a", + "shared_tun_secret_a", + Some(&first_peer_ns.name), + None, + "10.144.250.2/24", + )); + let mut second_peer = runtime.instance(test_config( + "shared_tun_second_peer", + "shared_tun_network_b", + "shared_tun_secret_b", + Some(&second_peer_ns.name), + None, + "10.144.251.2/24", + )); + + first.run().await.unwrap(); + second.run().await.unwrap(); + first_peer.run().await.unwrap(); + second_peer.run().await.unwrap(); + + wait_tun_ready(&first, &dev_name).await; + wait_tun_ready(&second, &dev_name).await; + assert_eq!( + first.get_global_ctx().get_tun_device_name(), + second.get_global_ctx().get_tun_device_name() + ); + + first_peer.add_connector_url(first.ring_listener_url()); + second_peer.add_connector_url(second.ring_listener_url()); + + wait_for_condition( + || async { + first + .get_core_instance() + .route_snapshots() + .await + .iter() + .any(|route| route.peer_id == first_peer.peer_id()) + && second + .get_core_instance() + .route_snapshots() + .await + .iter() + .any(|route| route.peer_id == second_peer.peer_id()) + }, + WAIT, + ) + .await; + wait_for_condition( + || async { ping_test(&first_peer_ns.name, "10.144.250.1", None).await }, + WAIT, + ) + .await; + wait_for_condition( + || async { ping_test(&second_peer_ns.name, "10.144.251.1", None).await }, + WAIT, + ) + .await; + + drop_insts(vec![first, second, first_peer, second_peer]).await; +} + +#[cfg(feature = "proxy-cidr-monitor")] +#[tokio::test] +#[serial_test::serial] +async fn runtime_proxy_patch_adds_and_removes_os_route() { + use crate::proto::api::config::ConfigPatchAction; + + let lab = ProxyLab::new(); + let source_dev = test_dev_name(); + let destination_dev = test_dev_name(); + let runtime = SharedTestRuntime::new(); + let mut source = runtime.instance(test_config( + "shared_patch_source", + "shared_patch_network", + "shared_patch_secret", + Some(&lab.source_ns), + Some(&source_dev), + "10.144.244.1/24", + )); + let mut destination = runtime.instance(test_config( + "shared_patch_destination", + "shared_patch_network", + "shared_patch_secret", + Some(&lab.owner_ns), + Some(&destination_dev), + "10.144.244.2/24", + )); + + source.run().await.unwrap(); + destination.run().await.unwrap(); + wait_tun_ready(&source, &source_dev).await; + wait_tun_ready(&destination, &destination_dev).await; + destination.add_connector_url(source.ring_listener_url()); + wait_for_condition( + || async { + source + .get_core_instance() + .route_snapshots() + .await + .iter() + .any(|route| route.peer_id == destination.peer_id()) + }, + WAIT, + ) + .await; + assert!(!ipv4_route_exists_in_ns( + &lab.source_ns, + &format!("{PROXY_CIDR} dev {source_dev}") + )); + + patch_proxy_cidr(&destination, ConfigPatchAction::Add).await; + wait_proxy_route(&source, destination.peer_id(), PROXY_CIDR).await; + wait_for_condition( + || async { + ipv4_route_exists_in_ns(&lab.source_ns, &format!("{PROXY_CIDR} dev {source_dev}")) + }, + WAIT, + ) + .await; + + patch_proxy_cidr(&destination, ConfigPatchAction::Remove).await; + wait_proxy_route_absent(&source, destination.peer_id(), PROXY_CIDR).await; + wait_for_condition( + || async { + !ipv4_route_exists_in_ns(&lab.source_ns, &format!("{PROXY_CIDR} dev {source_dev}")) + }, + WAIT, + ) + .await; + + drop_insts(vec![source, destination]).await; +} + +#[cfg(feature = "magic-dns")] +#[tokio::test] +#[serial_test::serial] +async fn magic_dns_route_lives_until_last_shared_owner_leaves() { + use crate::instance::dns_server::MAGIC_DNS_FAKE_IP; + + let netns = TestNetnsGuard::new(short_name("svd"), "10.232.1.2/24", "fd42::2/64"); + let dev_name = test_dev_name(); + let runtime = SharedTestRuntime::new(); + let first_config = test_config( + "shared_dns_first", + "shared_dns_network", + "shared_dns_secret", + Some(&netns.name), + Some(&dev_name), + "10.144.243.1/24", + ); + let mut flags = first_config.get_flags(); + flags.accept_dns = true; + first_config.set_flags(flags.clone()); + let second_config = test_config( + "shared_dns_second", + "shared_dns_network", + "shared_dns_secret", + Some(&netns.name), + Some(&dev_name), + "10.144.242.2/24", + ); + second_config.set_flags(flags); + let mut first = runtime.instance(first_config); + let mut second = runtime.instance(second_config); + + first.run().await.unwrap(); + second.run().await.unwrap(); + wait_tun_ready(&first, &dev_name).await; + wait_tun_ready(&second, &dev_name).await; + + let route = SharedIpv4Route::new(MAGIC_DNS_FAKE_IP.parse::().unwrap(), 32, None); + wait_for_condition( + || async { shared_route_owner_count(&runtime.registry, &dev_name, &route).await == 2 }, + WAIT, + ) + .await; + assert!(ipv4_route_exists_in_ns( + &netns.name, + &format!("{MAGIC_DNS_FAKE_IP} dev {dev_name}") + )); + + drop_insts(vec![first]).await; + wait_for_condition( + || async { + shared_route_owner_count(&runtime.registry, &dev_name, &route).await == 1 + && ipv4_route_exists_in_ns( + &netns.name, + &format!("{MAGIC_DNS_FAKE_IP} dev {dev_name}"), + ) + }, + WAIT, + ) + .await; + + drop_insts(vec![second]).await; + wait_for_condition( + || async { + shared_route_owner_count(&runtime.registry, &dev_name, &route).await == 0 + && !ipv4_route_exists_in_ns( + &netns.name, + &format!("{MAGIC_DNS_FAKE_IP} dev {dev_name}"), + ) + }, + WAIT, + ) + .await; +} diff --git a/easytier/src/tests/three_node.rs b/easytier/src/tests/three_node.rs index 5505345b..15ecb5cf 100644 --- a/easytier/src/tests/three_node.rs +++ b/easytier/src/tests/three_node.rs @@ -1017,6 +1017,157 @@ pub async fn public_ipv6_auto_addr_reconnect_reuses_same_address() { drop_insts(vec![provider, client]).await; } +#[cfg(feature = "tun")] +#[tokio::test] +#[serial_test::serial] +pub async fn shared_tun_public_ipv6_auto_addr_end_to_end() { + let lab = PublicIpv6Lab::setup_with_topology(PublicIpv6LabTopology::DelegatedPrefix); + let provider_dev = format!("st{:08x}", rand::random::()); + let client_dev = format!("st{:08x}", rand::random::()); + let process_runtime = CoreProcessRuntime::new(); + let shared_virtual_nic_registry = Instance::new_shared_virtual_nic_registry(); + + let provider_cfg = get_public_ipv6_config( + "provider_shared_public_ipv6", + PublicIpv6Lab::PROVIDER_NS, + "10.144.144.1", + &provider_dev, + uuid::Uuid::parse_str("44444444-4444-4444-4444-444444444444").unwrap(), + ); + provider_cfg.set_ipv6_public_addr_provider(true); + + let client_cfg = get_public_ipv6_config( + "client_shared_public_ipv6", + PublicIpv6Lab::CLIENT_NS, + "10.144.144.2", + &client_dev, + uuid::Uuid::parse_str("55555555-5555-5555-5555-555555555555").unwrap(), + ); + client_cfg.set_ipv6_public_addr_auto(true); + + let client_peer_cfg = get_public_ipv6_config( + "client_shared_public_ipv6_peer", + PublicIpv6Lab::CLIENT_NS, + "10.144.145.3", + &client_dev, + uuid::Uuid::parse_str("66666666-6666-6666-6666-666666666666").unwrap(), + ); + client_peer_cfg.set_listeners(vec![]); + + let mut provider = Instance::new_with_process_runtime_and_shared_virtual_nic_registry( + provider_cfg, + process_runtime.clone(), + shared_virtual_nic_registry.clone(), + ); + let mut client = Instance::new_with_process_runtime_and_shared_virtual_nic_registry( + client_cfg, + process_runtime.clone(), + shared_virtual_nic_registry.clone(), + ); + let mut client_peer = Instance::new_with_process_runtime_and_shared_virtual_nic_registry( + client_peer_cfg, + process_runtime, + shared_virtual_nic_registry, + ); + let mut client_events = client.get_global_ctx().subscribe(); + let mut client_peer_events = client_peer.get_global_ctx().subscribe(); + + provider.run().await.unwrap(); + client.run().await.unwrap(); + client_peer.run().await.unwrap(); + + let shared_ifname = wait_for_tun_ready(&mut client_events).await; + assert_eq!( + shared_ifname, + wait_for_tun_ready(&mut client_peer_events).await + ); + assert_eq!(shared_ifname, client_dev); + + provider.add_connector_url("tcp://10.1.1.2:11010".parse().unwrap()); + + wait_for_condition( + || async { + provider.get_core_instance().route_snapshots().await.len() == 1 + && client.get_core_instance().route_snapshots().await.len() == 1 + }, + Duration::from_secs(8), + ) + .await; + + wait_for_condition( + || async { + provider + .get_core_instance() + .node_snapshot() + .await + .ipv6_public_addr_prefix + == Some(PublicIpv6Lab::PROVIDER_PREFIX.parse().unwrap()) + }, + Duration::from_secs(10), + ) + .await; + + let leased = wait_for_public_ipv6_addr(&client).await; + wait_for_public_ipv6_route(&provider, leased).await; + + wait_for_condition( + || async { + addr_exists_in_ns(PublicIpv6Lab::CLIENT_NS, &client_dev, &leased.to_string()) + && route_exists_in_ns( + PublicIpv6Lab::CLIENT_NS, + &format!("default dev {client_dev}"), + ) + && route_exists_in_ns( + PublicIpv6Lab::PROVIDER_NS, + &format!("{} dev {provider_dev}", leased.address()), + ) + }, + Duration::from_secs(10), + ) + .await; + + wait_for_condition( + || async { ping6_test(PublicIpv6Lab::CLIENT_NS, PublicIpv6Lab::SERVER_IP, None).await }, + Duration::from_secs(10), + ) + .await; + + wait_for_condition( + || async { + ping6_test( + PublicIpv6Lab::SERVER_NS, + leased.address().to_string().as_str(), + None, + ) + .await + }, + Duration::from_secs(10), + ) + .await; + + drop_insts(vec![provider, client, client_peer]).await; + drop(lab); +} + +#[cfg(feature = "tun")] +async fn wait_for_tun_ready( + receiver: &mut tokio::sync::broadcast::Receiver, +) -> String { + tokio::time::timeout(Duration::from_secs(5), async { + loop { + match receiver.recv().await.unwrap() { + crate::common::global_ctx::GlobalCtxEvent::TunDeviceReady(ifname) => return ifname, + crate::common::global_ctx::GlobalCtxEvent::TunDeviceError(error) => { + panic!("tun device error: {error}") + } + _ => {} + } + } + }) + .await + .expect("timed out waiting for tun ready") +} + #[rstest::rstest] #[tokio::test] #[serial_test::serial]