From 800c840bb568b26c637be4028f2aa9e42820dbf4 Mon Sep 17 00:00:00 2001 From: "sijie.sun" Date: Tue, 16 Jun 2026 00:15:35 +0800 Subject: [PATCH] shared-tun: preserve member ownership on mobile Introduce shared NIC source ownership and dispatcher handling so a single dev_name can be shared by multiple tun-enabled instances while keeping per-member IP and route claims distinct. Pass Android VpnService fd registration with per-instance source and route claims. Keep the VPN address list limited to real member addresses and allow AF_INET6 without installing hidden fd00::1. Invalidate dispatcher flow and NAT state when source ownership changes or a member unregisters. Avoid rewriting non-first IPv4 fragment payloads, and adjust fragmented TCP/UDP checksums without recomputing over partial fragment bodies. Preserve source-owner routing for equal-prefix route conflicts, keep ICMP echo NAT entries distinct by echo id, and retry stale flow-owner send failures from the original packet. Only record NAT state after a translated packet is accepted by its member. Apply Linux IPv4 route preferred-source hints for shared routes and keep route repair paths source-aware. Keep Darwin ifcfg access scoped to cleanup-only paths where netns is not available. --- easytier-gui/src-tauri/src/lib.rs | 98 +- easytier-gui/src/composables/backend.ts | 21 +- easytier-gui/src/composables/mobile_vpn.ts | 124 +- easytier/src/common/ifcfg/darwin.rs | 18 +- easytier/src/common/ifcfg/mod.rs | 20 + easytier/src/common/ifcfg/netlink.rs | 220 ++- easytier/src/instance/instance.rs | 8 +- easytier/src/instance/shared_virtual_nic.rs | 385 ++++- .../instance/shared_virtual_nic/dispatcher.rs | 1513 ++++++++++++++++- easytier/src/instance/virtual_nic.rs | 149 +- easytier/src/instance_manager.rs | 22 + easytier/src/launcher.rs | 17 +- .../android/src/main/java/TauriVpnService.kt | 3 +- 13 files changed, 2423 insertions(+), 175 deletions(-) diff --git a/easytier-gui/src-tauri/src/lib.rs b/easytier-gui/src-tauri/src/lib.rs index 25fd6da0..fc2f28c5 100644 --- a/easytier-gui/src-tauri/src/lib.rs +++ b/easytier-gui/src-tauri/src/lib.rs @@ -66,6 +66,9 @@ static RPC_SERVER: once_cell::sync::Lazy>> = static WEB_CLIENT: once_cell::sync::Lazy>> = once_cell::sync::Lazy::new(|| RwLock::new(None)); +#[cfg(target_os = "android")] +const ANDROID_SHARED_TUN_DEV_NAME: &str = "easytier-shared"; + macro_rules! get_client_manager { () => {{ let guard = CLIENT_MANAGER @@ -76,6 +79,17 @@ macro_rules! get_client_manager { }}; } +fn normalize_network_config_for_runtime(cfg: &mut NetworkConfig) { + #[cfg(target_os = "android")] + { + if !cfg.no_tun() && cfg.dev_name.as_deref().map(str::is_empty).unwrap_or(true) { + cfg.dev_name = Some(ANDROID_SHARED_TUN_DEV_NAME.to_owned()); + } + } + #[cfg(not(target_os = "android"))] + let _ = cfg; +} + #[tauri::command] fn easytier_version() -> Result { Ok(easytier::VERSION.to_string()) @@ -100,6 +114,8 @@ fn set_dock_visibility(app: tauri::AppHandle, visible: bool) -> Result<(), Strin #[tauri::command] fn parse_network_config(cfg: NetworkConfig) -> Result { + let mut cfg = cfg; + normalize_network_config_for_runtime(&mut cfg); let toml = cfg.gen_config().map_err(|e| e.to_string())?; Ok(toml.dump()) } @@ -117,6 +133,8 @@ async fn run_network_instance( cfg: NetworkConfig, save: bool, ) -> Result<(), String> { + let mut cfg = cfg; + normalize_network_config_for_runtime(&mut cfg); let client_manager = get_client_manager!()?; let toml_config = cfg.gen_config().map_err(|e| e.to_string())?; client_manager @@ -155,16 +173,84 @@ async fn set_logging_level(level: String) -> Result<(), String> { Ok(()) } +#[derive(serde::Deserialize)] +#[serde(rename_all = "camelCase")] +#[allow(dead_code)] +struct TunFdInstanceSources { + instance_id: String, + ipv4_addrs: Vec, + #[serde(default)] + ipv6_addrs: Vec, + #[serde(default)] + ipv4_routes: Vec, + #[serde(default)] + ipv6_routes: Vec, +} + +#[cfg(mobile)] +fn parse_tun_fd_instance_sources( + instance_sources: Vec, +) -> Result< + std::collections::HashMap, + String, +> { + instance_sources + .into_iter() + .map(|source| { + let instance_id = source + .instance_id + .parse::() + .map_err(|err| format!("invalid instance id {}: {err}", source.instance_id))?; + let sources = easytier::instance::virtual_nic::MobileTunSources::parse( + source.ipv4_addrs, + source.ipv6_addrs, + source.ipv4_routes, + source.ipv6_routes, + ) + .map_err(|err| err.to_string())?; + Ok((instance_id, sources)) + }) + .collect() +} + #[tauri::command] -async fn set_tun_fd(fd: i32) -> Result<(), String> { +async fn set_tun_fd( + fd: i32, + instance_ids: Option>, + instance_sources: Option>, +) -> Result<(), String> { let Some(instance_manager) = INSTANCE_MANAGER.read().await.clone() else { return Err("set_tun_fd is not supported in remote mode".to_string()); }; + let target_ids = match instance_ids { + Some(instance_ids) if !instance_ids.is_empty() => instance_ids + .into_iter() + .map(|id| { + id.parse::() + .map_err(|err| format!("invalid instance id {id}: {err}")) + }) + .collect::, _>>()?, + _ => get_client_manager!()?.get_enabled_instances_for_tun_fd(), + }; + + #[cfg(mobile)] + let mut source_map = parse_tun_fd_instance_sources(instance_sources.unwrap_or_default())?; + #[cfg(not(mobile))] + let _ = instance_sources; + let mut success_count = 0; let mut errors = Vec::new(); - for uuid in get_client_manager!()?.get_enabled_instances_for_tun_fd() { - match instance_manager.set_tun_fd(&uuid, fd) { + for uuid in target_ids { + #[cfg(mobile)] + let set_result = match source_map.remove(&uuid) { + Some(sources) => instance_manager.set_tun_fd(&uuid, fd, sources), + None => Err(anyhow::anyhow!("missing tun sources for instance {uuid}")), + }; + #[cfg(not(mobile))] + let set_result = instance_manager.set_tun_fd(&uuid, fd); + + match set_result { Ok(()) => { success_count += 1; } @@ -1210,10 +1296,12 @@ mod manager { ) -> anyhow::Result<()> { self.storage.network_configs.clear(); for stored in configs { - let instance_id = stored.config.instance_id(); + let mut config = stored.config; + normalize_network_config_for_runtime(&mut config); + let instance_id = config.instance_id(); self.storage.network_configs.insert( instance_id.parse()?, - GUIConfig::new(instance_id.to_string(), stored.config, stored.source), + GUIConfig::new(instance_id.to_string(), config, stored.source), ); } diff --git a/easytier-gui/src/composables/backend.ts b/easytier-gui/src/composables/backend.ts index a16835f5..ebaf0c5d 100644 --- a/easytier-gui/src/composables/backend.ts +++ b/easytier-gui/src/composables/backend.ts @@ -70,8 +70,23 @@ export async function setLoggingLevel(level: string) { return await invoke('set_logging_level', { level }) } -export async function setTunFd(fd: number) { - return await invoke('set_tun_fd', { fd }) +export interface TunFdInstanceSources { + instanceId: string + ipv4Addrs: string[] + ipv6Addrs?: string[] + ipv4Routes?: string[] + ipv6Routes?: string[] +} + +export async function setTunFd(fd: number, instanceIds?: string[], instanceSources?: TunFdInstanceSources[]) { + const args: { fd: number, instanceIds?: string[], instanceSources?: TunFdInstanceSources[] } = { fd } + if (instanceIds?.length) { + args.instanceIds = instanceIds + } + if (instanceSources?.length) { + args.instanceSources = instanceSources + } + return await invoke('set_tun_fd', args) } export async function getEasytierVersion() { @@ -110,7 +125,7 @@ export async function sendConfigs(enabledNetworks: string[]) { config: NetworkTypes.toBackendNetworkConfig(config), source, })), - enabledNetworks + enabledNetworks, }) } diff --git a/easytier-gui/src/composables/mobile_vpn.ts b/easytier-gui/src/composables/mobile_vpn.ts index 3d3e2b01..ed0d0cab 100644 --- a/easytier-gui/src/composables/mobile_vpn.ts +++ b/easytier-gui/src/composables/mobile_vpn.ts @@ -2,6 +2,7 @@ import type { NetworkTypes } from 'easytier-frontend-lib' import { addPluginListener } from '@tauri-apps/api/core' import { Utils } from 'easytier-frontend-lib' import { get_vpn_status, prepare_vpn, start_vpn, stop_vpn } from 'tauri-plugin-vpnservice-api' +import type { TunFdInstanceSources } from './backend' type Route = NetworkTypes.Route @@ -12,6 +13,8 @@ interface vpnStatus { ipv4Cidr: number | null | undefined routes: string[] dns: string | null | undefined + instanceIds: string[] + instanceSources: TunFdInstanceSources[] } let dhcpPollingTimer: NodeJS.Timeout | null = null @@ -26,6 +29,8 @@ const curVpnStatus: vpnStatus = { ipv4Cidr: undefined, routes: [], dns: undefined, + instanceIds: [], + instanceSources: [], } async function requestVpnPermission() { @@ -50,6 +55,8 @@ function resetVpnConfigStatus() { curVpnStatus.ipv4Cidr = undefined curVpnStatus.routes = [] curVpnStatus.dns = undefined + curVpnStatus.instanceIds = [] + curVpnStatus.instanceSources = [] } function syncVpnStatusFromNative(status: Awaited>) { @@ -59,13 +66,14 @@ function syncVpnStatusFromNative(status: Awaited { + await setTunFd(payload.fd, curVpnStatus.instanceIds, curVpnStatus.instanceSources).catch((e) => { console.error('set tun fd failed', e) }) } @@ -218,12 +236,73 @@ function getRoutesForVpn(routes: Route[], node_config: NetworkTypes.NetworkConfi return Array.from(new Set(ret)).sort() } +function ipv4CidrToRoute(cidr: string): string | undefined { + const [address, prefixText] = cidr.split('/') + const prefix = Number(prefixText) + const octets = address?.split('.').map(octet => Number(octet)) + + if ( + octets?.length !== 4 + || !Number.isInteger(prefix) + || prefix < 0 + || prefix > 32 + || octets.some((octet) => { + return !Number.isInteger(octet) || octet < 0 || octet > 255 + }) + ) { + return undefined + } + + const ip = ( + octets[0] * 0x1000000 + + octets[1] * 0x10000 + + octets[2] * 0x100 + + octets[3] + ) >>> 0 + const mask = prefix === 0 ? 0 : (0xFFFFFFFF << (32 - prefix)) >>> 0 + const network = (ip & mask) >>> 0 + const route = [ + (network >>> 24) & 0xFF, + (network >>> 16) & 0xFF, + (network >>> 8) & 0xFF, + network & 0xFF, + ].join('.') + + return `${route}/${prefix}` +} + function getCollectedNetworkInfo(response: Awaited>, instanceId: string) { const info = response.info as any const map = info?.map ?? info return map?.[instanceId] } +function sortInstanceSources(sources: TunFdInstanceSources[]): TunFdInstanceSources[] { + return sources + .map(source => ({ + instanceId: source.instanceId, + ipv4Addrs: [...source.ipv4Addrs].sort(), + ipv6Addrs: [...(source.ipv6Addrs ?? [])].sort(), + ipv4Routes: [...(source.ipv4Routes ?? [])].sort(), + ipv6Routes: [...(source.ipv6Routes ?? [])].sort(), + })) + .sort((a, b) => a.instanceId.localeCompare(b.instanceId)) +} + +function splitRoutesByFamily(routes: string[]) { + const ipv4Routes: string[] = [] + const ipv6Routes: string[] = [] + routes.forEach((route) => { + if (route.includes(':')) { + ipv6Routes.push(route) + } + else { + ipv4Routes.push(route) + } + }) + return { ipv4Routes, ipv6Routes } +} + export async function onNetworkInstanceChange(instanceId: string) { pendingVpnConfigInstanceId = instanceId if (!vpnConfigSyncTask) { @@ -273,6 +352,7 @@ async function applyNetworkInstanceChange(instanceId: string) { } const ipv4Addrs: string[] = [] + const instanceSources: TunFdInstanceSources[] = [] const routes = new Set() let dns: string | undefined const retryInstanceIds: string[] = [] @@ -308,8 +388,26 @@ async function applyNetworkInstanceChange(instanceId: string) { network_length = 24 } - ipv4Addrs.push(`${virtual_ip}/${network_length}`) - getRoutesForVpn(curNetworkInfo?.routes, config).forEach(route => routes.add(route)) + const sourceIpv4 = `${virtual_ip}/${network_length}` + ipv4Addrs.push(sourceIpv4) + const instanceRoutes = new Set() + const localRoute = ipv4CidrToRoute(sourceIpv4) + if (localRoute) { + routes.add(localRoute) + instanceRoutes.add(localRoute) + } + getRoutesForVpn(curNetworkInfo?.routes, config).forEach((route) => { + routes.add(route) + instanceRoutes.add(route) + }) + const { ipv4Routes, ipv6Routes } = splitRoutesByFamily([...instanceRoutes]) + instanceSources.push({ + instanceId, + ipv4Addrs: [sourceIpv4], + ipv6Addrs: [], + ipv4Routes, + ipv6Routes, + }) if (config.enable_magic_dns) { dns = '100.100.100.101' } @@ -330,10 +428,14 @@ async function applyNetworkInstanceChange(instanceId: string) { const sortedIpv4Addrs = [...ipv4Addrs].sort() const sortedRoutes = Array.from(routes).sort() + const sortedInstanceIds = group.map(({ instanceId }) => instanceId).sort() + const sortedInstanceSources = sortInstanceSources(instanceSources) const ipChanged = JSON.stringify(sortedIpv4Addrs) !== JSON.stringify(curVpnStatus.ipv4Addrs) const routesChanged = JSON.stringify(sortedRoutes) !== JSON.stringify(curVpnStatus.routes) - const dnsChanged = dns != curVpnStatus.dns - const configChanged = ipChanged || routesChanged || dnsChanged + const dnsChanged = dns !== curVpnStatus.dns + const instanceIdsChanged = JSON.stringify(sortedInstanceIds) !== JSON.stringify(curVpnStatus.instanceIds) + const instanceSourcesChanged = JSON.stringify(sortedInstanceSources) !== JSON.stringify(sortInstanceSources(curVpnStatus.instanceSources)) + const configChanged = ipChanged || routesChanged || dnsChanged || instanceIdsChanged || instanceSourcesChanged const shouldStartVpn = !curVpnStatus.running if (shouldStartVpn || configChanged) { @@ -357,7 +459,7 @@ async function applyNetworkInstanceChange(instanceId: string) { } try { - await doStartVpn(sortedIpv4Addrs, sortedRoutes, dns) + await doStartVpn(sortedIpv4Addrs, sortedRoutes, dns, sortedInstanceIds, sortedInstanceSources) } catch (e) { if (e instanceof Error && e.message === 'need_prepare') { diff --git a/easytier/src/common/ifcfg/darwin.rs b/easytier/src/common/ifcfg/darwin.rs index fe8cec98..34dbe0d0 100644 --- a/easytier/src/common/ifcfg/darwin.rs +++ b/easytier/src/common/ifcfg/darwin.rs @@ -46,12 +46,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(), diff --git a/easytier/src/common/ifcfg/mod.rs b/easytier/src/common/ifcfg/mod.rs index 7f97744c..db62789f 100644 --- a/easytier/src/common/ifcfg/mod.rs +++ b/easytier/src/common/ifcfg/mod.rs @@ -31,6 +31,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, @@ -39,6 +49,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, diff --git a/easytier/src/common/ifcfg/netlink.rs b/easytier/src/common/ifcfg/netlink.rs index 13c0ec65..14431ec7 100644 --- a/easytier/src/common/ifcfg/netlink.rs +++ b/easytier/src/common/ifcfg/netlink.rs @@ -376,6 +376,64 @@ impl NetlinkIfConfiger { pub(crate) fn list_ipv6_route_messages() -> Result, Error> { Self::list_route_messages(AddressFamily::Inet6) } + + fn ipv4_route_message( + ifindex: u32, + address: Ipv4Addr, + cidr_prefix: u8, + cost: Option, + source_hint: Option, + ) -> RouteMessage { + let mut message = RouteMessage::default(); + + message.header.table = RouteHeader::RT_TABLE_MAIN; + message.header.protocol = RouteProtocol::Static; + message.header.scope = RouteScope::Universe; + message.header.kind = RouteType::Unicast; + message.header.address_family = AddressFamily::Inet; + message.header.destination_prefix_length = cidr_prefix; + + message + .attributes + .push(RouteAttribute::Priority(cost.unwrap_or(65535) as u32)); + message.attributes.push(RouteAttribute::Oif(ifindex)); + message + .attributes + .push(RouteAttribute::Destination(RouteAddress::Inet(address))); + + if let Some(source_hint) = source_hint { + message + .attributes + .push(RouteAttribute::PrefSource(RouteAddress::Inet(source_hint))); + } + + message + } + + fn ipv4_route_target_matches( + route: &Route, + address: Ipv4Addr, + cidr_prefix: u8, + ifidx: u32, + ) -> bool { + route.destination == IpAddr::V4(address) + && route.prefix == cidr_prefix + && route.ifindex == Some(ifidx) + } + + fn ipv4_route_exact_matches( + route: &Route, + address: Ipv4Addr, + cidr_prefix: u8, + ifidx: u32, + cost: Option, + source_hint: Option, + ) -> bool { + Self::ipv4_route_target_matches(route, address, cidr_prefix, ifidx) + && route.table == RouteHeader::RT_TABLE_MAIN + && route.metric == Some(cost.unwrap_or(65535) as u32) + && route.source_hint == source_hint.map(IpAddr::V4) + } } #[async_trait] @@ -387,29 +445,25 @@ impl IfConfiguerTrait for NetlinkIfConfiger { cidr_prefix: u8, cost: Option, ) -> Result<(), Error> { - let mut message = RouteMessage::default(); - - message.header.table = RouteHeader::RT_TABLE_MAIN; - message.header.protocol = RouteProtocol::Static; - message.header.scope = RouteScope::Universe; - message.header.kind = RouteType::Unicast; - message.header.address_family = AddressFamily::Inet; - // metric - message - .attributes - .push(RouteAttribute::Priority(cost.unwrap_or(65535) as u32)); - // output interface - message - .attributes - .push(RouteAttribute::Oif(NetlinkIfConfiger::get_interface_index( - name, - )?)); - // source address - message.header.destination_prefix_length = cidr_prefix; - message - .attributes - .push(RouteAttribute::Destination(RouteAddress::Inet(address))); + 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, + ); send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::NewRoute(message), false) } @@ -424,10 +478,41 @@ impl IfConfiguerTrait for NetlinkIfConfiger { for msg in routes { let other_route: Route = msg.clone().into(); - if other_route.destination == std::net::IpAddr::V4(address) - && other_route.prefix == cidr_prefix - && other_route.ifindex == Some(ifidx) - { + if NetlinkIfConfiger::ipv4_route_target_matches( + &other_route, + address, + cidr_prefix, + ifidx, + ) { + send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::DelRoute(msg), true)?; + return Ok(()); + } + } + + 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 routes = Self::list_routes()?; + let ifidx = NetlinkIfConfiger::get_interface_index(name)?; + + for msg in routes { + let other_route: Route = msg.clone().into(); + if NetlinkIfConfiger::ipv4_route_exact_matches( + &other_route, + address, + cidr_prefix, + ifidx, + cost, + source_hint, + ) { send_netlink_req_and_wait_one_resp(RouteNetlinkMessage::DelRoute(msg), true)?; return Ok(()); } @@ -666,6 +751,89 @@ mod tests { } } + #[test] + fn ipv4_route_message_includes_pref_source_when_source_hint_is_set() { + let source_hint = Ipv4Addr::new(10, 231, 1, 1); + let message = NetlinkIfConfiger::ipv4_route_message( + 7, + Ipv4Addr::new(10, 99, 0, 0), + 24, + Some(123), + Some(source_hint), + ); + + assert_eq!(message.header.destination_prefix_length, 24); + assert!(message.attributes.iter().any(|attr| { + matches!( + attr, + RouteAttribute::PrefSource(RouteAddress::Inet(source)) if *source == source_hint + ) + })); + assert!(message.attributes.iter().any(|attr| { + matches!(attr, RouteAttribute::Priority(priority) if *priority == 123) + })); + assert!( + message + .attributes + .iter() + .any(|attr| matches!(attr, RouteAttribute::Oif(7))) + ); + } + + #[test] + fn ipv4_route_exact_match_distinguishes_metric_and_pref_source() { + let address = Ipv4Addr::new(10, 99, 0, 0); + let source_hint = Ipv4Addr::new(10, 99, 0, 1); + let other_source_hint = Ipv4Addr::new(10, 99, 0, 2); + let route: Route = + NetlinkIfConfiger::ipv4_route_message(7, address, 24, Some(123), Some(source_hint)) + .into(); + + assert!(NetlinkIfConfiger::ipv4_route_exact_matches( + &route, + address, + 24, + 7, + Some(123), + Some(source_hint), + )); + assert!(!NetlinkIfConfiger::ipv4_route_exact_matches( + &route, + address, + 24, + 7, + Some(124), + Some(source_hint), + )); + assert!(!NetlinkIfConfiger::ipv4_route_exact_matches( + &route, + address, + 24, + 7, + Some(123), + Some(other_source_hint), + )); + assert!(!NetlinkIfConfiger::ipv4_route_exact_matches( + &route, + address, + 24, + 7, + Some(123), + None, + )); + + let mut non_main_table_route = route.clone(); + non_main_table_route.table = 100; + assert!(!NetlinkIfConfiger::ipv4_route_exact_matches( + &non_main_table_route, + address, + 24, + 7, + Some(123), + Some(source_hint), + )); + } + struct PrepareEnv {} impl PrepareEnv { fn new() -> Self { diff --git a/easytier/src/instance/instance.rs b/easytier/src/instance/instance.rs index d94713cf..c0cbbf74 100644 --- a/easytier/src/instance/instance.rs +++ b/easytier/src/instance/instance.rs @@ -193,6 +193,8 @@ impl IpProxy { #[cfg(feature = "tun")] type NicCtx = super::virtual_nic::NicCtx; +#[cfg(all(feature = "tun", mobile))] +use super::virtual_nic::MobileTunSources; #[cfg(feature = "magic-dns")] struct MagicDnsContainer { @@ -910,7 +912,8 @@ impl Instance { close_notifier: Arc, shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, ) -> Result { - if global_ctx.get_flags().dev_name.is_empty() { + let flags = global_ctx.get_flags(); + if flags.dev_name.is_empty() { return Ok(NicCtx::new( global_ctx, peer_manager, @@ -1730,6 +1733,7 @@ impl Instance { peer_packet_receiver: Arc>, shared_virtual_nic_registry: ArcSharedVirtualNicRegistry, fd: i32, + sources: MobileTunSources, ) -> Result<(), anyhow::Error> { tracing::info!("setup_nic_ctx_for_mobile, fd: {}", fd); Self::clear_nic_ctx(nic_ctx.clone(), peer_packet_receiver.clone()).await; @@ -1747,7 +1751,7 @@ impl Instance { .await .with_context(|| "create nic ctx failed")?; new_nic_ctx - .run_for_mobile(fd) + .run_for_mobile(fd, sources) .await .with_context(|| "add ip failed")?; diff --git a/easytier/src/instance/shared_virtual_nic.rs b/easytier/src/instance/shared_virtual_nic.rs index 1dc7170c..07c1aa4e 100644 --- a/easytier/src/instance/shared_virtual_nic.rs +++ b/easytier/src/instance/shared_virtual_nic.rs @@ -96,6 +96,9 @@ pub struct SharedIfConfigDelta { pub ipv4_addresses: OwnedItemDelta, pub ipv6_addresses: OwnedItemDelta, pub ipv4_routes: OwnedItemDelta, + pub ipv4_route_removed_old_source_hints: BTreeMap>, + pub ipv4_route_source_changed: BTreeSet, + pub ipv4_route_source_changed_old_hints: BTreeMap>, pub ipv6_routes: OwnedItemDelta, pub mtu: Option, } @@ -131,6 +134,12 @@ impl SharedIfConfig { .cloned() .unwrap_or_default(); let old_mtu = self.effective_mtu(); + let source_change_candidates = old_claims + .ipv4_routes + .union(&claims.ipv4_routes) + .cloned() + .collect::>(); + let old_ipv4_route_sources = self.ipv4_route_sources(&source_change_candidates); let ipv4_addresses = update_owned_items( &mut self.ipv4_address_owners, @@ -159,11 +168,22 @@ impl SharedIfConfig { update_member_mtu(&mut self.member_mtu, member_id, claims.mtu); self.member_claims.insert(member_id, claims); + let ipv4_route_removed_old_source_hints = + old_ipv4_route_hints(&old_ipv4_route_sources, &ipv4_routes.removed); + let ipv4_route_source_changed_old_hints = + self.changed_ipv4_route_sources(&old_ipv4_route_sources, &ipv4_routes); + let ipv4_route_source_changed = ipv4_route_source_changed_old_hints + .keys() + .cloned() + .collect(); SharedIfConfigDelta { ipv4_addresses, ipv6_addresses, ipv4_routes, + ipv4_route_removed_old_source_hints, + ipv4_route_source_changed, + ipv4_route_source_changed_old_hints, ipv6_routes, mtu: mtu_delta(old_mtu, self.effective_mtu()), } @@ -173,8 +193,10 @@ impl SharedIfConfig { &mut self, member_id: SharedVirtualNicMemberId, ) -> Option { - let old_claims = self.member_claims.remove(&member_id)?; + let old_claims = self.member_claims.get(&member_id).cloned()?; let old_mtu = self.effective_mtu(); + let old_ipv4_route_sources = self.ipv4_route_sources(&old_claims.ipv4_routes); + self.member_claims.remove(&member_id); let ipv4_addresses = remove_owned_items( &mut self.ipv4_address_owners, @@ -198,11 +220,22 @@ impl SharedIfConfig { ); self.member_mtu.remove(&member_id); + let ipv4_route_removed_old_source_hints = + old_ipv4_route_hints(&old_ipv4_route_sources, &ipv4_routes.removed); + let ipv4_route_source_changed_old_hints = + self.changed_ipv4_route_sources(&old_ipv4_route_sources, &ipv4_routes); + let ipv4_route_source_changed = ipv4_route_source_changed_old_hints + .keys() + .cloned() + .collect(); Some(SharedIfConfigDelta { ipv4_addresses, ipv6_addresses, ipv4_routes, + ipv4_route_removed_old_source_hints, + ipv4_route_source_changed, + ipv4_route_source_changed_old_hints, ipv6_routes, mtu: mtu_delta(old_mtu, self.effective_mtu()), }) @@ -250,6 +283,57 @@ impl SharedIfConfig { .cloned() .unwrap_or_default() } + + fn ipv4_route_source_hint(&self, route: &SharedIpv4Route) -> Option { + let owners = self.ipv4_route_owners.get(route)?; + let route_inet = Ipv4Inet::new(route.address, route.prefix).ok(); + let mut fallback = None; + + for owner in owners { + let Some(claims) = self.member_claims.get(owner) else { + 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 ipv4_route_sources( + &self, + routes: &BTreeSet, + ) -> BTreeMap> { + routes + .iter() + .map(|route| (route.clone(), self.ipv4_route_source_hint(route))) + .collect() + } + + fn changed_ipv4_route_sources( + &self, + old_sources: &BTreeMap>, + route_delta: &OwnedItemDelta, + ) -> BTreeMap> { + old_sources + .iter() + .filter(|(route, old_source)| { + !route_delta.added.contains(*route) + && !route_delta.removed.contains(*route) + && self.ipv4_route_owners.contains_key(*route) + && self.ipv4_route_source_hint(route) != **old_source + }) + .map(|(route, old_source)| (route.clone(), *old_source)) + .collect() + } } pub struct SharedVirtualNic { @@ -509,7 +593,24 @@ impl SharedVirtualNic { let nic = self.nic.lock().await; for route in &delta.ipv4_routes.removed { - ignore_removed_ifcfg_not_found(nic.remove_route(route.address, route.prefix).await)?; + let source_hint = delta + .ipv4_route_removed_old_source_hints + .get(route) + .copied() + .flatten(); + ignore_removed_ifcfg_not_found( + remove_shared_ipv4_route(&nic, route, source_hint).await, + )?; + } + for route in &delta.ipv4_route_source_changed { + let source_hint = delta + .ipv4_route_source_changed_old_hints + .get(route) + .copied() + .flatten(); + 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( @@ -531,8 +632,10 @@ impl SharedVirtualNic { .await?; } for route in &delta.ipv4_routes.added { - nic.add_route_with_cost(route.address, route.prefix, route.cost) - .await?; + add_shared_ipv4_route(&nic, route, _next_ifcfg).await?; + } + for route in &delta.ipv4_route_source_changed { + 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) @@ -548,8 +651,7 @@ impl SharedVirtualNic { if !delta.ipv4_addresses.removed.is_empty() { for route in _next_ifcfg.ipv4_route_owners.keys() { ignore_added_ifcfg_already_exists( - nic.add_route_with_cost(route.address, route.prefix, route.cost) - .await, + add_shared_ipv4_route(&nic, route, _next_ifcfg).await, )?; } } @@ -631,9 +733,7 @@ impl SharedVirtualNic { dispatcher: &SharedVirtualNicDispatcher, ) -> Result<(), Error> { for (member_id, claims) in &self.ifcfg.member_claims { - dispatcher - .update_sources(*member_id, &claims.ipv4_addresses, &claims.ipv6_addresses) - .await?; + dispatcher.update_sources(*member_id, claims).await?; } Ok(()) } @@ -644,17 +744,9 @@ impl SharedVirtualNic { old_claims: &SharedIfConfigClaims, next_claims: &SharedIfConfigClaims, ) -> Result<(), Error> { - let mut active_ipv4_addresses = old_claims.ipv4_addresses.clone(); - active_ipv4_addresses.extend(next_claims.ipv4_addresses.iter().copied()); - let mut active_ipv6_addresses = old_claims.ipv6_addresses.clone(); - active_ipv6_addresses.extend(next_claims.ipv6_addresses.iter().copied()); - - self.sync_dispatcher_sources_for_addresses( - member_id, - &active_ipv4_addresses, - &active_ipv6_addresses, - ) - .await + 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_member( @@ -662,24 +754,17 @@ impl SharedVirtualNic { member_id: SharedVirtualNicMemberId, claims: &SharedIfConfigClaims, ) -> Result<(), Error> { - self.sync_dispatcher_sources_for_addresses( - member_id, - &claims.ipv4_addresses, - &claims.ipv6_addresses, - ) - .await + self.sync_dispatcher_sources_for_claims(member_id, claims) + .await } - async fn sync_dispatcher_sources_for_addresses( + async fn sync_dispatcher_sources_for_claims( &self, member_id: SharedVirtualNicMemberId, - ipv4_addresses: &BTreeSet, - ipv6_addresses: &BTreeSet, + claims: &SharedIfConfigClaims, ) -> Result<(), Error> { if let Some(dispatcher) = &self.dispatcher { - dispatcher - .update_sources(member_id, ipv4_addresses, ipv6_addresses) - .await?; + dispatcher.update_sources(member_id, claims).await?; } Ok(()) } @@ -695,6 +780,75 @@ impl SharedVirtualNic { } } +fn dispatcher_claims_for_ifcfg_transition( + old_claims: &SharedIfConfigClaims, + next_claims: &SharedIfConfigClaims, +) -> SharedIfConfigClaims { + let mut claims = SharedIfConfigClaims::default(); + claims + .ipv4_addresses + .extend(old_claims.ipv4_addresses.iter().copied()); + claims + .ipv4_addresses + .extend(next_claims.ipv4_addresses.iter().copied()); + claims + .ipv6_addresses + .extend(old_claims.ipv6_addresses.iter().copied()); + claims + .ipv6_addresses + .extend(next_claims.ipv6_addresses.iter().copied()); + claims + .ipv4_routes + .extend(old_claims.ipv4_routes.iter().cloned()); + claims + .ipv4_routes + .extend(next_claims.ipv4_routes.iter().cloned()); + claims + .ipv6_routes + .extend(old_claims.ipv6_routes.iter().cloned()); + claims + .ipv6_routes + .extend(next_claims.ipv6_routes.iter().cloned()); + claims +} + +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 old_ipv4_route_hints( + old_sources: &BTreeMap>, + routes: &BTreeSet, +) -> BTreeMap> { + routes + .iter() + .filter_map(|route| { + old_sources + .get(route) + .map(|source_hint| (route.clone(), *source_hint)) + }) + .collect() +} + fn ignore_removed_ifcfg_not_found(result: Result<(), Error>) -> Result<(), Error> { match result { Err(Error::NotFound) => Ok(()), @@ -897,8 +1051,7 @@ impl SharedVirtualNicMember { } #[cfg(mobile)] - pub async fn add_mobile_source_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { - let ip = ipv4_inet(ip, cidr)?; + pub async fn add_mobile_source_ip(&self, ip: Ipv4Inet) -> Result<(), Error> { self.update_claims_for_mobile(|claims| { claims.ipv4_addresses.insert(ip); }) @@ -906,14 +1059,29 @@ impl SharedVirtualNicMember { } #[cfg(mobile)] - pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { - let ip = ipv6_inet(ip, cidr)?; + pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Inet) -> Result<(), Error> { self.update_claims_for_mobile(|claims| { claims.ipv6_addresses.insert(ip); }) .await } + #[cfg(mobile)] + pub async fn add_mobile_source_ipv4_route(&self, route: SharedIpv4Route) -> Result<(), Error> { + self.update_claims_for_mobile(|claims| { + claims.ipv4_routes.insert(route); + }) + .await + } + + #[cfg(mobile)] + pub async fn add_mobile_source_ipv6_route(&self, route: SharedIpv6Route) -> Result<(), Error> { + self.update_claims_for_mobile(|claims| { + claims.ipv6_routes.insert(route); + }) + .await + } + pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { self.update_claims(|claims| match ip { Some(ip) => { @@ -1292,6 +1460,17 @@ mod tests { } } + 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)) } @@ -1385,6 +1564,140 @@ mod tests { ); } + #[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 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!( + delta.ipv4_route_source_changed, + BTreeSet::from([route.clone()]) + ); + assert_eq!( + delta.ipv4_route_source_changed_old_hints, + BTreeMap::from([(route.clone(), 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 delta = ifcfg.remove_member(second).unwrap(); + + assert!(delta.ipv4_routes.removed.is_empty()); + assert_eq!( + delta.ipv4_route_source_changed, + BTreeSet::from([route.clone()]) + ); + assert_eq!( + delta.ipv4_route_source_changed_old_hints, + BTreeMap::from([(route.clone(), 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_records_old_source_hint_per_cost() { + 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 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!( + delta.ipv4_route_removed_old_source_hints, + BTreeMap::from([(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()); diff --git a/easytier/src/instance/shared_virtual_nic/dispatcher.rs b/easytier/src/instance/shared_virtual_nic/dispatcher.rs index b4f7f666..51032711 100644 --- a/easytier/src/instance/shared_virtual_nic/dispatcher.rs +++ b/easytier/src/instance/shared_virtual_nic/dispatcher.rs @@ -10,6 +10,14 @@ use std::{ use cidr::{Ipv4Inet, Ipv6Inet}; use futures::{SinkExt, StreamExt}; +use pnet::packet::{ + MutablePacket as _, Packet as _, + icmp::{self, MutableIcmpPacket}, + ip::IpNextHeaderProtocols, + ipv4::{self, MutableIpv4Packet}, + tcp::{self, MutableTcpPacket}, + udp::{self, MutableUdpPacket}, +}; #[cfg(mobile)] use std::sync::OnceLock; #[cfg(mobile)] @@ -26,7 +34,10 @@ use crate::{ tunnel::{Tunnel, ZCPacketSink, ZCPacketStream, packet_def::ZCPacket}, }; -use super::{SharedVirtualNicMemberId, SharedVirtualNicMemberRegistrationId}; +use super::{ + SharedIfConfigClaims, SharedIpv4Route, SharedIpv6Route, SharedVirtualNicMemberId, + SharedVirtualNicMemberRegistrationId, +}; const MEMBER_TUNNEL_BUFFER_SIZE: usize = 1024; const FLOW_OWNER_LIMIT: usize = 4096; @@ -34,6 +45,8 @@ 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 TCP_PROTOCOL: u8 = 6; const UDP_PROTOCOL: u8 = 17; const DISPATCHER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(1); @@ -71,7 +84,7 @@ enum SharedVirtualNicControl { }, UpdateSources { member_id: SharedVirtualNicMemberId, - sources: BTreeSet, + sources: SharedVirtualNicMemberSources, ack: oneshot::Sender<()>, }, Shutdown { @@ -250,6 +263,95 @@ struct SharedVirtualNicFlowKey { ports: Option, } +#[derive(Clone, Debug, Default, PartialEq, Eq)] +struct SharedVirtualNicMemberSources { + exact: BTreeSet, + 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 { + exact: ipv4_addresses + .iter() + .map(|addr| SharedVirtualNicFlowAddr::from(addr.address())) + .chain( + ipv6_addresses + .iter() + .map(|addr| SharedVirtualNicFlowAddr::from(addr.address())), + ) + .collect(), + ipv4_addresses, + ipv6_addresses, + ipv4_routes, + ipv6_routes, + } + } + + fn is_empty(&self) -> bool { + self.exact.is_empty() + && self.ipv4_addresses.is_empty() + && self.ipv6_addresses.is_empty() + && self.ipv4_routes.is_empty() + && self.ipv6_routes.is_empty() + } +} + +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, + } + } +} + impl SharedVirtualNicFlowKey { fn from_packet(packet: &ZCPacket) -> Option { let payload = packet.payload(); @@ -272,13 +374,18 @@ impl SharedVirtualNicFlowKey { } 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: transport_ports(protocol, &payload[header_len..]), + ports: if fragment_offset == 0 { + transport_ports(protocol, &payload[header_len..]) + } else { + None + }, }) } @@ -301,7 +408,13 @@ impl SharedVirtualNicFlowKey { src: self.dst, dst: self.src, protocol: self.protocol, - ports: self.ports.map(|ports| ports.reversed()), + ports: self.ports.map(|ports| { + if self.protocol == ICMP_PROTOCOL { + ports + } else { + ports.reversed() + } + }), } } } @@ -331,12 +444,6 @@ impl SharedVirtualNicFlowTable { self.owners.get(&key).copied() } - fn remove_owner(&mut self, member_id: SharedVirtualNicMemberId) { - self.owners.retain(|_, owner| *owner != member_id); - self.insert_order - .retain(|key| self.owners.contains_key(key)); - } - fn clear(&mut self) { self.owners.clear(); self.insert_order.clear(); @@ -353,6 +460,71 @@ impl SharedVirtualNicFlowTable { } } +#[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, @@ -435,14 +607,13 @@ impl SharedVirtualNicDispatcher { pub(super) async fn update_sources( &self, member_id: SharedVirtualNicMemberId, - ipv4_addresses: &BTreeSet, - ipv6_addresses: &BTreeSet, + claims: &SharedIfConfigClaims, ) -> Result<(), Error> { let (ack, rx) = oneshot::channel(); self.control_sender .send(SharedVirtualNicControl::UpdateSources { member_id, - sources: sources_from_addresses(ipv4_addresses, ipv6_addresses), + sources: SharedVirtualNicMemberSources::from_claims(claims), ack, }) .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; @@ -458,7 +629,7 @@ impl SharedVirtualNicDispatcher { self.control_sender .send(SharedVirtualNicControl::UpdateSources { member_id, - sources: BTreeSet::new(), + sources: SharedVirtualNicMemberSources::default(), ack, }) .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; @@ -554,9 +725,10 @@ impl SharedVirtualNicDispatcherTask { &mut self, member_packet: SharedVirtualNicMemberPacket, ) -> bool { - self.state - .remember_reverse_owner(member_packet.member_id, &member_packet.packet); - if let Err(err) = self.tun_sink.send(member_packet.packet).await { + 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; } @@ -708,8 +880,9 @@ impl SharedVirtualNicMobileDispatcherTask { &mut self, member_packet: SharedVirtualNicMemberPacket, ) -> bool { - self.state - .remember_reverse_owner(member_packet.member_id, &member_packet.packet); + 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, @@ -718,7 +891,7 @@ impl SharedVirtualNicMobileDispatcherTask { return false; }; - if let Err(err) = tun_sink.send(member_packet.packet).await { + if let Err(err) = tun_sink.send(packet).await { tracing::error!(?err, "shared virtual nic write to mobile tun failed"); self.drop_tun(); return true; @@ -769,6 +942,7 @@ fn next_mobile_rebuild_delay(delay: Duration) -> Duration { struct SharedVirtualNicDispatcherState { members: BTreeMap, flow_table: SharedVirtualNicFlowTable, + nat_table: SharedVirtualNicNatTable, source_table: SharedVirtualNicSourceTable, } @@ -789,6 +963,8 @@ impl SharedVirtualNicDispatcherState { sources, ack, } => { + self.flow_table.clear(); + self.nat_table.clear(); self.source_table.update_member_sources(member_id, sources); let _ = ack.send(()); } @@ -823,12 +999,15 @@ impl SharedVirtualNicDispatcherState { let entry = self.members.remove(&member_id); drop(entry); - self.flow_table.remove_owner(member_id); + self.flow_table.clear(); + 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.flow_table.clear(); + self.nat_table.clear(); self.source_table.clear(); for entry in members.into_values() { @@ -840,6 +1019,16 @@ impl SharedVirtualNicDispatcherState { self.flow_table.remember_reverse_owner(member_id, packet); } + fn prepare_member_packet_to_tun( + &mut self, + member_id: SharedVirtualNicMemberId, + mut packet: ZCPacket, + ) -> ZCPacket { + self.nat_table.translate_reply(&mut packet); + self.remember_reverse_owner(member_id, &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"); @@ -847,18 +1036,68 @@ impl SharedVirtualNicDispatcherState { } async fn send_packet(&mut self, packet: ZCPacket) -> bool { - let mut packet = packet; + let source_owner = self.source_table.owner_of_source(&packet, &self.members); + let preferred_destination_owner = source_owner.active_member(); - if let Some(member_id) = self.flow_table.owner_of(&packet) { - match self.send_packet_to_member(member_id, packet).await { + let flow_owner = self.flow_table.owner_of(&packet); + let destination_owner = self.source_table.owner_of_destination( + &packet, + &self.members, + preferred_destination_owner, + ); + let mut packet = packet; + if let Some(member_id) = flow_owner { + let original_packet = packet.clone(); + let result = if should_translate_source_for_member(source_owner, member_id) { + self.send_packet_to_member_with_translation(member_id, packet) + .await + } else { + self.send_packet_to_member(member_id, packet).await + }; + + match result { Ok(()) => return true, - Err(packet_on_failure) => { - packet = packet_on_failure; + Err(_) => { + self.flow_table.clear(); + self.nat_table.clear(); + packet = original_packet; } } + + 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, + ); + return self + .send_packet_without_flow_owner(packet, source_owner, destination_owner) + .await; + } + + self.send_packet_without_flow_owner(packet, source_owner, destination_owner) + .await + } + + async fn send_packet_without_flow_owner( + &mut self, + packet: ZCPacket, + source_owner: SourceOwner, + destination_owner: Option, + ) -> bool { + 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 self.source_table.owner_of_source(&packet, &self.members) { + match source_owner { SourceOwner::Active(member_id) => { return self.send_packet_to_member(member_id, packet).await.is_ok(); } @@ -866,11 +1105,34 @@ impl SharedVirtualNicDispatcherState { SourceOwner::None => {} } - let Some(member_id) = self.members.keys().next().copied() else { - return false; - }; + false + } - self.send_packet_to_member(member_id, packet).await.is_ok() + async fn send_packet_to_member_with_translation( + &mut self, + member_id: SharedVirtualNicMemberId, + mut packet: ZCPacket, + ) -> Result<(), ZCPacket> { + let mut nat_entry = None; + if let Some(key) = SharedVirtualNicFlowKey::from_packet(&packet) { + if let Some(translated_src) = self + .source_table + .source_for_member_destination(member_id, key.dst) + { + if key.src != translated_src + && rewrite_packet_source(&mut packet, key.src, translated_src) + { + nat_entry = Some((packet.clone(), key.src, translated_src)); + } + } + } + + self.send_packet_to_member(member_id, packet).await?; + 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( @@ -903,9 +1165,226 @@ enum SourceOwner { 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_ipv4_addr(packet, expected_source, new_source, RewriteIpv4Addr::Source) +} + +fn rewrite_packet_destination( + packet: &mut ZCPacket, + expected_destination: SharedVirtualNicFlowAddr, + new_destination: SharedVirtualNicFlowAddr, +) -> bool { + rewrite_ipv4_addr( + packet, + expected_destination, + new_destination, + RewriteIpv4Addr::Destination, + ) +} + +#[derive(Clone, Copy)] +enum RewriteIpv4Addr { + Source, + Destination, +} + +fn rewrite_ipv4_addr( + packet: &mut ZCPacket, + expected_addr: SharedVirtualNicFlowAddr, + new_addr: SharedVirtualNicFlowAddr, + rewrite: RewriteIpv4Addr, +) -> 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; + } + + let old_source = ipv4_packet.get_source(); + let old_destination = ipv4_packet.get_destination(); + let is_fragmented = ipv4_packet.get_fragment_offset() != 0 + || (ipv4_packet.get_flags() & ipv4::Ipv4Flags::MoreFragments) != 0; + + match rewrite { + RewriteIpv4Addr::Source if ipv4_packet.get_source() == expected_addr => { + ipv4_packet.set_source(new_addr); + } + RewriteIpv4Addr::Destination if ipv4_packet.get_destination() == expected_addr => { + ipv4_packet.set_destination(new_addr); + } + _ => return false, + } + + if !is_fragmented { + update_ipv4_transport_checksum(&mut ipv4_packet, header_len); + } else if ipv4_packet.get_fragment_offset() == 0 { + adjust_ipv4_fragment_transport_checksum( + &mut ipv4_packet, + header_len, + old_source, + old_destination, + ); + } + ipv4_packet.set_checksum(0); + let checksum = ipv4::checksum(&ipv4_packet.to_immutable()); + ipv4_packet.set_checksum(checksum); + true +} + +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_ipv4_fragment_transport_checksum( + ipv4_packet: &mut MutableIpv4Packet<'_>, + header_len: usize, + old_source: std::net::Ipv4Addr, + old_destination: std::net::Ipv4Addr, +) { + 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; + }; + let checksum = adjust_ipv4_pseudo_header_checksum( + tcp_packet.get_checksum(), + old_source, + source, + old_destination, + destination, + ); + tcp_packet.set_checksum(checksum); + } + IpNextHeaderProtocols::Udp => { + let Some(mut udp_packet) = MutableUdpPacket::new(transport_payload) else { + return; + }; + let checksum = udp_packet.get_checksum(); + if checksum == 0 { + return; + } + let checksum = adjust_ipv4_pseudo_header_checksum( + checksum, + old_source, + source, + old_destination, + destination, + ); + udp_packet.set_checksum(if checksum == 0 { 0xffff } else { checksum }); + } + _ => {} + } +} + +fn adjust_ipv4_pseudo_header_checksum( + checksum: u16, + old_source: std::net::Ipv4Addr, + source: std::net::Ipv4Addr, + old_destination: std::net::Ipv4Addr, + destination: std::net::Ipv4Addr, +) -> u16 { + let mut checksum = checksum; + for (old, new) in ipv4_checksum_words(old_source).zip(ipv4_checksum_words(source)) { + checksum = adjust_checksum_word(checksum, old, new); + } + for (old, new) in ipv4_checksum_words(old_destination).zip(ipv4_checksum_words(destination)) { + checksum = adjust_checksum_word(checksum, old, new); + } + checksum +} + +fn ipv4_checksum_words(addr: std::net::Ipv4Addr) -> impl Iterator { + let octets = addr.octets(); + [ + u16::from_be_bytes([octets[0], octets[1]]), + u16::from_be_bytes([octets[2], octets[3]]), + ] + .into_iter() +} + +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>, + member_sources: BTreeMap, source_owners: BTreeMap>, } @@ -913,14 +1392,14 @@ impl SharedVirtualNicSourceTable { fn update_member_sources( &mut self, member_id: SharedVirtualNicMemberId, - sources: BTreeSet, + sources: SharedVirtualNicMemberSources, ) { let old_sources = self.member_sources.remove(&member_id).unwrap_or_default(); - for source in old_sources.difference(&sources) { + for source in old_sources.exact.difference(&sources.exact) { self.remove_source_owner(*source, member_id); } - for source in sources.difference(&old_sources) { + for source in sources.exact.difference(&old_sources.exact) { self.source_owners .entry(*source) .or_default() @@ -937,7 +1416,7 @@ impl SharedVirtualNicSourceTable { return; }; - for source in sources { + for source in sources.exact { self.remove_source_owner(source, member_id); } } @@ -967,6 +1446,47 @@ impl SharedVirtualNicSourceTable { .unwrap_or(SourceOwner::Inactive) } + fn owner_of_destination( + &self, + packet: &ZCPacket, + active_members: &BTreeMap, + preferred_member: Option, + ) -> Option { + let dst = SharedVirtualNicFlowKey::from_packet(packet)?.dst; + + let mut best = None; + for (member_id, sources) in &self.member_sources { + if !active_members.contains_key(member_id) { + continue; + } + let Some(prefix) = sources.destination_prefix(dst) else { + continue; + }; + if best + .map(|(best_member, best_prefix)| { + prefix > best_prefix + || (prefix == best_prefix + && Some(*member_id) == preferred_member + && best_member != *member_id) + }) + .unwrap_or(true) + { + best = Some((*member_id, prefix)); + } + } + best.map(|(member_id, _)| member_id) + } + + fn source_for_member_destination( + &self, + member_id: SharedVirtualNicMemberId, + dst: SharedVirtualNicFlowAddr, + ) -> Option { + self.member_sources + .get(&member_id)? + .source_for_destination(dst) + } + fn remove_source_owner( &mut self, source: SharedVirtualNicFlowAddr, @@ -983,25 +1503,139 @@ impl SharedVirtualNicSourceTable { } } -fn sources_from_addresses( - ipv4_addresses: &BTreeSet, - ipv6_addresses: &BTreeSet, -) -> BTreeSet { - ipv4_addresses - .iter() - .map(|addr| SharedVirtualNicFlowAddr::V4(u32::from_be_bytes(addr.address().octets()))) - .chain( - ipv6_addresses - .iter() - .map(|addr| SharedVirtualNicFlowAddr::V6(addr.address().octets())), - ) - .collect() +impl SharedVirtualNicMemberSources { + fn destination_prefix(&self, dst: SharedVirtualNicFlowAddr) -> Option { + match dst { + SharedVirtualNicFlowAddr::V4(dst) => { + let dst = std::net::Ipv4Addr::from(dst); + self.ipv4_destination_prefix(dst) + } + SharedVirtualNicFlowAddr::V6(dst) => { + let dst = std::net::Ipv6Addr::from(dst); + self.ipv6_destination_prefix(dst) + } + } + } + + fn source_for_destination( + &self, + dst: SharedVirtualNicFlowAddr, + ) -> Option { + match dst { + SharedVirtualNicFlowAddr::V4(dst) => { + let dst = std::net::Ipv4Addr::from(dst); + self.ipv4_source_for_destination(dst) + .map(SharedVirtualNicFlowAddr::from) + } + SharedVirtualNicFlowAddr::V6(dst) => { + let dst = std::net::Ipv6Addr::from(dst); + self.ipv6_source_for_destination(dst) + .map(SharedVirtualNicFlowAddr::from) + } + } + } + + fn ipv4_destination_prefix(&self, dst: std::net::Ipv4Addr) -> Option { + self.ipv4_addresses + .iter() + .filter(|addr| addr.contains(&dst)) + .map(|addr| addr.network_length()) + .chain( + self.ipv4_routes + .iter() + .filter(|route| route.contains(&dst)) + .map(|route| route.network_length()), + ) + .max() + } + + fn ipv6_destination_prefix(&self, dst: std::net::Ipv6Addr) -> Option { + self.ipv6_addresses + .iter() + .filter(|addr| addr.contains(&dst)) + .map(|addr| addr.network_length()) + .chain( + self.ipv6_routes + .iter() + .filter(|route| route.contains(&dst)) + .map(|route| route.network_length()), + ) + .max() + } + + 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 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 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_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 => ICMP_ECHO_HEADER_LEN, _ => return None, }; @@ -1009,10 +1643,25 @@ fn transport_ports(protocol: u8, payload: &[u8]) -> Option icmp_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 read_ipv6_addr(payload: &[u8], start: usize) -> [u8; 16] { @@ -1023,7 +1672,10 @@ fn read_ipv6_addr(payload: &[u8], start: usize) -> [u8; 16] { #[cfg(test)] mod tests { - use std::{net::Ipv6Addr, time::Duration}; + use std::{ + net::{Ipv4Addr, Ipv6Addr}, + time::Duration, + }; use super::*; use crate::tunnel::{TunnelError, common::TunnelWrapper, ring::create_ring_tunnel_pair}; @@ -1037,6 +1689,186 @@ mod tests { 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 ipv4_non_first_fragment(src: Ipv4Addr, dst: Ipv4Addr, fragment_payload: &[u8]) -> ZCPacket { + let mut payload = vec![0; IPV4_HEADER_MIN_LEN + fragment_payload.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_fragment_offset(1); + ipv4_packet.set_source(src); + ipv4_packet.set_destination(dst); + } + payload[IPV4_HEADER_MIN_LEN..].copy_from_slice(fragment_payload); + { + 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_first_fragment_with_more_fragments( + src: Ipv4Addr, + dst: Ipv4Addr, + checksum: u16, + ) -> ZCPacket { + let mut payload = vec![0; IPV4_HEADER_MIN_LEN + UDP_HEADER_LEN + 4]; + 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_flags(ipv4::Ipv4Flags::MoreFragments); + 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(1234); + udp_packet.set_destination(5678); + udp_packet.set_length((UDP_HEADER_LEN + 8) as u16); + udp_packet.set_checksum(checksum); + } + payload[IPV4_HEADER_MIN_LEN + UDP_HEADER_LEN..].copy_from_slice(&[1, 2, 3, 4]); + { + 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_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)) } @@ -1057,7 +1889,6 @@ mod tests { fn source_table_selects_ipv6_source_owner() { let first = uuid::Uuid::from_u128(1); let second = uuid::Uuid::from_u128(2); - let first_addr = "2001:db8::1".parse::().unwrap(); let second_addr = "2001:db8::2".parse::().unwrap(); let dst = "2001:db8:ffff::1".parse::().unwrap(); let mut table = SharedVirtualNicSourceTable::default(); @@ -1067,14 +1898,8 @@ mod tests { members.insert(first, member_entry(first_sender)); members.insert(second, member_entry(second_sender)); - table.update_member_sources( - first, - BTreeSet::from([SharedVirtualNicFlowAddr::V6(first_addr.octets())]), - ); - table.update_member_sources( - second, - BTreeSet::from([SharedVirtualNicFlowAddr::V6(second_addr.octets())]), - ); + 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), @@ -1088,27 +1913,132 @@ mod tests { ); } + #[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_over_fallback_member() { - let fallback = uuid::Uuid::from_u128(1); + async fn dispatcher_prefers_source_owner_for_equal_prefix_route_conflict() { + 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(10, 99, 0, 2); + 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"], &["10.99.0.0/24"]), + ); + state.source_table.update_member_sources( + source_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(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_preserves_non_first_fragment_payload() { + let src = Ipv4Addr::new(10, 231, 1, 1); + let translated_src = Ipv4Addr::new(10, 231, 2, 1); + let dst = Ipv4Addr::new(10, 231, 2, 2); + let fragment_payload = [0x12, 0x34, 0x56, 0x78, 0x9a, 0xbc, 0xde, 0xf0]; + let mut packet = ipv4_non_first_fragment(src, dst, &fragment_payload); + + assert!(rewrite_packet_source( + &mut packet, + SharedVirtualNicFlowAddr::from(src), + SharedVirtualNicFlowAddr::from(translated_src) + )); + + let ipv4 = pnet::packet::ipv4::Ipv4Packet::new(packet.payload()).unwrap(); + assert_eq!(ipv4.get_source(), translated_src); + assert_eq!(ipv4.get_destination(), dst); + assert_eq!(&packet.payload()[IPV4_HEADER_MIN_LEN..], &fragment_payload); + let key = SharedVirtualNicFlowKey::from_packet(&packet).unwrap(); + assert_eq!(key.ports, None); + } + + #[test] + fn rewrite_ipv4_source_adjusts_first_fragment_transport_checksum() { + let src = Ipv4Addr::new(10, 231, 1, 1); + let translated_src = Ipv4Addr::new(10, 231, 2, 1); + let dst = Ipv4Addr::new(10, 231, 2, 2); + let checksum = 0x1234; + let mut packet = ipv4_udp_first_fragment_with_more_fragments(src, dst, checksum); + let expected_checksum = + adjust_ipv4_pseudo_header_checksum(checksum, src, translated_src, dst, dst); + + assert!(rewrite_packet_source( + &mut packet, + SharedVirtualNicFlowAddr::from(src), + SharedVirtualNicFlowAddr::from(translated_src) + )); + + let udp = + pnet::packet::udp::UdpPacket::new(&packet.payload()[IPV4_HEADER_MIN_LEN..]).unwrap(); + assert_eq!(udp.get_checksum(), expected_checksum); + assert_eq!( + &packet.payload()[IPV4_HEADER_MIN_LEN + UDP_HEADER_LEN..], + &[1, 2, 3, 4] + ); + } + + #[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 (fallback_sender, mut fallback_receiver) = mpsc::channel(1); + 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(fallback, member_entry(fallback_sender)); + state.register(first, member_entry(first_sender)); state.register(owner, member_entry(owner_sender)); - state.source_table.update_member_sources( - owner, - BTreeSet::from([SharedVirtualNicFlowAddr::V6(source.octets())]), - ); + 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!(fallback_receiver.try_recv().is_err()); + assert!(first_receiver.try_recv().is_err()); assert!(owner_receiver.try_recv().is_ok()); } @@ -1124,10 +2054,9 @@ mod tests { state.register(fallback, member_entry(fallback_sender)); state.register(owner, member_entry(owner_sender)); - state.source_table.update_member_sources( - owner, - BTreeSet::from([SharedVirtualNicFlowAddr::V6(source.octets())]), - ); + 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)) @@ -1136,6 +2065,429 @@ mod tests { 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_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_flow_owners_globally() { + let stale_owner = uuid::Uuid::from_u128(2); + let new_owner = uuid::Uuid::from_u128(1); + let claimed_ip = Ipv4Addr::new(10, 231, 1, 1); + let remote_ip = Ipv4Addr::new(203, 0, 113, 1); + let (stale_sender, mut stale_receiver) = mpsc::channel(1); + let (new_sender, mut new_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register(stale_owner, member_entry(stale_sender)); + state.register(new_owner, member_entry(new_sender)); + state + .source_table + .update_member_sources(stale_owner, member_sources(&["10.231.1.1/24"])); + state.remember_reverse_owner( + stale_owner, + &ipv4_udp_packet_with_ports(remote_ip, claimed_ip, 5678, 1234), + ); + + let (ack, _rx) = oneshot::channel(); + state.handle_control(SharedVirtualNicControl::UpdateSources { + member_id: new_owner, + sources: member_sources(&["10.231.1.1/24"]), + ack, + }); + state + .forward_tun_packet_to_member(ipv4_udp_packet_with_ports( + claimed_ip, remote_ip, 1234, 5678, + )) + .await; + + assert!(stale_receiver.try_recv().is_err()); + assert!(new_receiver.try_recv().is_ok()); + } + + #[tokio::test] + async fn dispatcher_flow_owner_failure_retries_original_packet() { + let stale_owner = uuid::Uuid::from_u128(1); + let source_owner = uuid::Uuid::from_u128(2); + let stale_owner_ip = Ipv4Addr::new(10, 231, 2, 1); + let source_owner_ip = Ipv4Addr::new(10, 231, 1, 1); + let remote_ip = Ipv4Addr::new(10, 231, 2, 2); + let (stale_sender, stale_receiver) = mpsc::channel(1); + let (source_sender, mut source_receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + drop(stale_receiver); + state.register(stale_owner, member_entry(stale_sender)); + state.register(source_owner, member_entry(source_sender)); + state + .source_table + .update_member_sources(stale_owner, member_sources(&["10.231.2.1/24"])); + state + .source_table + .update_member_sources(source_owner, member_sources(&["10.231.1.1/24"])); + state.remember_reverse_owner( + stale_owner, + &ipv4_udp_packet_with_ports(remote_ip, source_owner_ip, 5678, 1234), + ); + + state + .forward_tun_packet_to_member(ipv4_udp_packet_with_ports( + source_owner_ip, + remote_ip, + 1234, + 5678, + )) + .await; + + 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(), stale_owner_ip); + assert_eq!(ipv4.get_destination(), remote_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); @@ -1150,6 +2502,9 @@ mod tests { 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)) @@ -1185,8 +2540,9 @@ mod tests { close_notifier.clone(), ) .unwrap(); + let empty_claims = SharedIfConfigClaims::default(); dispatcher - .update_sources(member_id, &BTreeSet::new(), &BTreeSet::new()) + .update_sources(member_id, &empty_claims) .await .unwrap(); @@ -1226,8 +2582,9 @@ mod tests { close_notifier.clone(), ) .unwrap(); + let empty_claims = SharedIfConfigClaims::default(); dispatcher - .update_sources(member_id, &BTreeSet::new(), &BTreeSet::new()) + .update_sources(member_id, &empty_claims) .await .unwrap(); diff --git a/easytier/src/instance/virtual_nic.rs b/easytier/src/instance/virtual_nic.rs index 3fb0997b..3b38e775 100644 --- a/easytier/src/instance/virtual_nic.rs +++ b/easytier/src/instance/virtual_nic.rs @@ -46,10 +46,82 @@ use crate::common::ifcfg::RegistryManager; #[cfg(test)] use super::shared_virtual_nic::SharedVirtualNic; +#[cfg(mobile)] +use super::shared_virtual_nic::{SharedIpv4Route, SharedIpv6Route}; use super::shared_virtual_nic::{ SharedVirtualNicMember, SharedVirtualNicMemberId, SharedVirtualNicRegistry, }; +#[cfg(mobile)] +#[derive(Clone, Debug, Default)] +pub struct MobileTunSources { + pub ipv4: Vec, + pub ipv6: Vec, + pub ipv4_routes: Vec, + pub ipv6_routes: Vec, +} + +#[cfg(mobile)] +impl MobileTunSources { + pub fn parse( + ipv4: Vec, + ipv6: Vec, + ipv4_routes: Vec, + ipv6_routes: Vec, + ) -> Result { + let ipv4 = ipv4 + .into_iter() + .map(|addr| { + addr.parse::() + .map_err(|err| anyhow::anyhow!("invalid IPv4 source {addr}: {err}").into()) + }) + .collect::, Error>>()?; + let ipv6 = ipv6 + .into_iter() + .map(|addr| { + addr.parse::() + .map_err(|err| anyhow::anyhow!("invalid IPv6 source {addr}: {err}").into()) + }) + .collect::, Error>>()?; + let ipv4_routes = ipv4_routes + .into_iter() + .map(|route| { + let route = route.parse::().map_err(|err| { + let err: Error = + anyhow::anyhow!("invalid IPv4 route source {route}: {err}").into(); + err + })?; + Ok(SharedIpv4Route::new( + route.address(), + route.network_length(), + None, + )) + }) + .collect::, Error>>()?; + let ipv6_routes = ipv6_routes + .into_iter() + .map(|route| { + let route = route.parse::().map_err(|err| { + let err: Error = + anyhow::anyhow!("invalid IPv6 route source {route}: {err}").into(); + err + })?; + Ok(SharedIpv6Route::new( + route.address(), + route.network_length(), + None, + )) + }) + .collect::, Error>>()?; + Ok(Self { + ipv4, + ipv6, + ipv4_routes, + ipv6_routes, + }) + } +} + pin_project! { pub struct TunStream { #[pin] @@ -765,10 +837,21 @@ impl VirtualNic { 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, cost) + .add_ipv4_route_with_source_hint(self.ifname(), address, cidr, cost, source_hint) .await?; Ok(()) } @@ -781,6 +864,26 @@ impl VirtualNic { 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(()) + } + pub async fn add_ipv6_route(&self, address: Ipv6Addr, cidr: u8) -> Result<(), Error> { self.add_ipv6_route_with_cost(address, cidr, None).await } @@ -1031,18 +1134,34 @@ impl NicBackend { } #[cfg(mobile)] - pub async fn add_mobile_source_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { + pub async fn add_mobile_source_ip(&self, ip: Ipv4Inet) -> Result<(), Error> { match self { Self::Dedicated(_) => Ok(()), - Self::Shared(member) => member.add_mobile_source_ip(ip, cidr).await, + Self::Shared(member) => member.add_mobile_source_ip(ip).await, } } #[cfg(mobile)] - pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { + pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Inet) -> Result<(), Error> { match self { Self::Dedicated(_) => Ok(()), - Self::Shared(member) => member.add_mobile_source_ipv6(ip, cidr).await, + Self::Shared(member) => member.add_mobile_source_ipv6(ip).await, + } + } + + #[cfg(mobile)] + pub async fn add_mobile_source_ipv4_route(&self, route: SharedIpv4Route) -> Result<(), Error> { + match self { + Self::Dedicated(_) => Ok(()), + Self::Shared(member) => member.add_mobile_source_ipv4_route(route).await, + } + } + + #[cfg(mobile)] + pub async fn add_mobile_source_ipv6_route(&self, route: SharedIpv6Route) -> Result<(), Error> { + match self { + Self::Dedicated(_) => Ok(()), + Self::Shared(member) => member.add_mobile_source_ipv6_route(route).await, } } } @@ -1739,7 +1858,11 @@ impl NicCtx { } #[cfg(mobile)] - pub async fn run_for_mobile(&mut self, tun_fd: std::os::fd::RawFd) -> Result<(), Error> { + pub async fn run_for_mobile( + &mut self, + tun_fd: std::os::fd::RawFd, + sources: MobileTunSources, + ) -> Result<(), Error> { let (tunnel, ifname) = match self.backend.create_dev_for_mobile(tun_fd).await { Ok(ret) => { let ifname = self @@ -1756,14 +1879,20 @@ impl NicCtx { } }; - if let Some(ipv4_addr) = self.global_ctx.get_ipv4() { + for ipv4_addr in sources.ipv4 { + self.backend.add_mobile_source_ip(ipv4_addr).await?; + } + for ipv6_addr in sources.ipv6 { + self.backend.add_mobile_source_ipv6(ipv6_addr).await?; + } + for ipv4_route in sources.ipv4_routes { self.backend - .add_mobile_source_ip(ipv4_addr.address(), ipv4_addr.network_length() as i32) + .add_mobile_source_ipv4_route(ipv4_route) .await?; } - if let Some(ipv6_addr) = self.global_ctx.get_ipv6() { + for ipv6_route in sources.ipv6_routes { self.backend - .add_mobile_source_ipv6(ipv6_addr.address(), ipv6_addr.network_length() as i32) + .add_mobile_source_ipv6_route(ipv6_route) .await?; } diff --git a/easytier/src/instance_manager.rs b/easytier/src/instance_manager.rs index a28f310a..495c2543 100644 --- a/easytier/src/instance_manager.rs +++ b/easytier/src/instance_manager.rs @@ -302,6 +302,28 @@ impl NetworkInstanceManager { .and_then(|instance| instance.value().get_api_service()) } + #[cfg(mobile)] + pub fn set_tun_fd( + &self, + instance_id: &uuid::Uuid, + fd: i32, + sources: crate::instance::virtual_nic::MobileTunSources, + ) -> Result<(), anyhow::Error> { + let sender = self + .instance_map + .get(instance_id) + .ok_or_else(|| anyhow::anyhow!("instance not found"))? + .get_tun_fd_sender() + .ok_or_else(|| anyhow::anyhow!("tun fd sender not found"))?; + + sender + .try_send(Some(crate::launcher::MobileTunFd { fd, sources })) + .map_err(|e| anyhow::anyhow!("failed to send tun fd: {}", e))?; + + Ok(()) + } + + #[cfg(not(mobile))] pub fn set_tun_fd(&self, instance_id: &uuid::Uuid, fd: i32) -> Result<(), anyhow::Error> { let sender = self .instance_map diff --git a/easytier/src/launcher.rs b/easytier/src/launcher.rs index e425a924..9be43fef 100644 --- a/easytier/src/launcher.rs +++ b/easytier/src/launcher.rs @@ -6,6 +6,8 @@ use crate::common::config::{ use crate::gateway::socks5::Socks5Server; #[cfg(feature = "ffi-dataplane")] pub use crate::gateway::socks5::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket}; +#[cfg(mobile)] +use crate::instance::virtual_nic::MobileTunSources; use crate::proto::api::{self, manage}; use crate::proto::rpc_types::controller::BaseController; use crate::rpc_service::InstanceRpcService; @@ -36,6 +38,16 @@ use tokio::{ pub type MyNodeInfo = crate::proto::api::manage::MyNodeInfo; type ArcMutApiService = Arc>>>; +#[cfg(mobile)] +#[derive(Clone, Debug)] +pub struct MobileTunFd { + pub fd: i32, + pub sources: MobileTunSources, +} + +#[cfg(mobile)] +type TunFd = Option; +#[cfg(not(mobile))] type TunFd = Option; #[derive(serde::Serialize, Clone)] @@ -135,11 +147,12 @@ impl EasyTierLauncher { peer_mgr.clone(), peer_packet_receiver.clone(), shared_virtual_nic_registry.clone(), - tun_fd, + tun_fd.fd, + tun_fd.sources, ) .await { - tracing::error!(?err, tun_fd, "setup mobile nic ctx failed"); + tracing::error!(?err, fd = tun_fd.fd, "setup mobile nic ctx failed"); } } }); diff --git a/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt b/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt index 41b77e5f..a19637c3 100644 --- a/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt +++ b/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt @@ -5,6 +5,7 @@ import android.net.VpnService import android.os.Build import android.os.ParcelFileDescriptor import android.os.Bundle +import android.system.OsConstants.AF_INET6 import java.net.InetAddress import java.util.Arrays @@ -115,7 +116,7 @@ class TauriVpnService : VpnService() { if (ipParts.size != 2) throw IllegalArgumentException("Invalid IP addr string") builder.addAddress(ipParts[0], ipParts[1].toInt()) } - builder.addAddress("fd00::1", 128) + builder.allowFamily(AF_INET6) builder.setMtu(mtu) dns?.let { builder.addDnsServer(it) }