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]