diff --git a/easytier-gui/src-tauri/src/lib.rs b/easytier-gui/src-tauri/src/lib.rs index 876f201f..b2338cdc 100644 --- a/easytier-gui/src-tauri/src/lib.rs +++ b/easytier-gui/src-tauri/src/lib.rs @@ -160,14 +160,31 @@ async fn set_tun_fd(fd: i32) -> 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()); }; - if let Some(uuid) = get_client_manager!()? - .get_enabled_instances_with_tun_ids() - .next() - { - instance_manager - .set_tun_fd(&uuid, fd) - .map_err(|e| e.to_string())?; + + 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) { + Ok(()) => { + success_count += 1; + } + Err(err) => { + errors.push(format!("{}: {}", uuid, err)); + } + } } + + if success_count == 0 && !errors.is_empty() { + return Err(format!( + "failed to set tun fd for all instances: {}", + errors.join("; ") + )); + } + + for err in errors { + eprintln!("set_tun_fd skipped instance: {err}"); + } + Ok(()) } @@ -918,32 +935,98 @@ mod manager { .filter_map(|c| c.config.instance_id().parse::().ok()) } + pub fn get_enabled_instances_for_tun_fd(&self) -> Vec { + let Some(first) = self + .storage + .network_configs + .iter() + .filter(|v| self.storage.enabled_networks.contains(v.key())) + .filter(|v| !v.config.no_tun()) + .find_map(|c| { + c.config + .instance_id() + .parse::() + .ok() + .map(|id| (id, Self::shared_tun_dev_name(&c.config).map(str::to_owned))) + }) + else { + return Vec::new(); + }; + + let (first_id, Some(shared_dev_name)) = first else { + return vec![first.0]; + }; + + let ids: Vec = self + .storage + .network_configs + .iter() + .filter(|v| self.storage.enabled_networks.contains(v.key())) + .filter(|v| Self::shared_tun_dev_name(&v.config) == Some(shared_dev_name.as_str())) + .filter_map(|c| c.config.instance_id().parse::().ok()) + .collect(); + if ids.is_empty() { vec![first_id] } else { ids } + } + + fn shared_tun_dev_name(config: &NetworkConfig) -> Option<&str> { + if config.no_tun() { + return None; + } + config + .dev_name + .as_deref() + .filter(|dev_name| !dev_name.is_empty()) + } + #[cfg(target_os = "android")] - pub fn get_enabled_instances_with_web_like_tun_ids( + fn runtime_shared_tun_dev_name( + cfg: &easytier::common::config::TomlConfigLoader, + ) -> Option { + let flags = cfg.get_flags(); + if flags.no_tun || flags.dev_name.is_empty() { + None + } else { + Some(flags.dev_name) + } + } + + #[cfg(target_os = "android")] + fn is_compatible_android_tun( + config: &NetworkConfig, + shared_dev_name: Option<&str>, + ) -> bool { + matches!( + (Self::shared_tun_dev_name(config), shared_dev_name), + (Some(existing), Some(next)) if existing == next + ) + } + + #[cfg(target_os = "android")] + fn enabled_incompatible_tun_ids( &self, - ) -> impl Iterator + '_ { + web_only: bool, + shared_dev_name: Option<&str>, + ) -> Vec { self.storage .network_configs .iter() .filter(|v| self.storage.enabled_networks.contains(v.key())) .filter(|v| !v.config.no_tun()) - .filter(|v| v.source.is_web_like()) + .filter(|v| !web_only || v.source.is_web_like()) + .filter(|v| !Self::is_compatible_android_tun(&v.config, shared_dev_name)) .filter_map(|c| c.config.instance_id().parse::().ok()) + .collect() } #[cfg(target_os = "android")] - pub(super) async fn disable_instances_with_tun( + pub(super) async fn disable_incompatible_instances_with_tun( &self, app: &AppHandle, web_only: bool, + shared_dev_name: Option<&str>, ) -> Result<(), easytier::rpc_service::remote_client::RemoteClientError> { - let inst_ids: Vec = if web_only { - self.get_enabled_instances_with_web_like_tun_ids().collect() - } else { - self.get_enabled_instances_with_tun_ids().collect() - }; - for inst_id in inst_ids { + for inst_id in self.enabled_incompatible_tun_ids(web_only, shared_dev_name) { self.handle_update_network_state(app.clone(), inst_id, true) .await?; } @@ -951,11 +1034,20 @@ mod manager { } pub(super) fn notify_vpn_stop_if_no_tun(&self, app: &AppHandle) -> Result<(), String> { - let has_tun = self.get_enabled_instances_with_tun_ids().any(|_| true); - if !has_tun { - app.emit("vpn_service_stop", "") + #[cfg(target_os = "android")] + if let Some(instance_id) = self.get_enabled_instances_with_tun_ids().next() { + app.emit("vpn_service_config_changed", instance_id.to_string()) .map_err(|e| e.to_string())?; + return Ok(()); } + + #[cfg(not(target_os = "android"))] + if self.get_enabled_instances_with_tun_ids().next().is_some() { + return Ok(()); + } + + app.emit("vpn_service_stop", "") + .map_err(|e| e.to_string())?; Ok(()) } @@ -971,19 +1063,31 @@ mod manager { #[cfg(target_os = "android")] if !cfg.get_flags().no_tun { + let shared_dev_name = Self::runtime_shared_tun_dev_name(cfg); match source { PersistedConfigSource::User | PersistedConfigSource::Legacy => { - self.disable_instances_with_tun(app, false) - .await - .map_err(|e| e.to_string())?; + self.disable_incompatible_instances_with_tun( + app, + false, + shared_dev_name.as_deref(), + ) + .await + .map_err(|e| e.to_string())?; } PersistedConfigSource::Web => { - self.disable_instances_with_tun(app, true) - .await - .map_err(|e| e.to_string())?; - if self.get_enabled_instances_with_tun_ids().next().is_some() { + self.disable_incompatible_instances_with_tun( + app, + true, + shared_dev_name.as_deref(), + ) + .await + .map_err(|e| e.to_string())?; + if !self + .enabled_incompatible_tun_ids(false, shared_dev_name.as_deref()) + .is_empty() + { return Err( - "Android only supports one active TUN network; user-managed VPN remains active" + "Android only supports one active TUN device; user-managed VPN remains active with an incompatible dev_name" .to_string(), ); } diff --git a/easytier-gui/src/auto-imports.d.ts b/easytier-gui/src/auto-imports.d.ts index 25d72636..b246d2dd 100644 --- a/easytier-gui/src/auto-imports.d.ts +++ b/easytier-gui/src/auto-imports.d.ts @@ -52,6 +52,7 @@ declare global { const mapWritableState: typeof import('pinia')['mapWritableState'] const markRaw: typeof import('vue')['markRaw'] const nextTick: typeof import('vue')['nextTick'] + const normalizeConfigSource: typeof import('./composables/config_source')['normalizeConfigSource'] const onActivated: typeof import('vue')['onActivated'] const onBeforeMount: typeof import('vue')['onBeforeMount'] const onBeforeRouteLeave: typeof import('vue-router')['onBeforeRouteLeave'] @@ -177,6 +178,7 @@ declare module 'vue' { readonly mapWritableState: UnwrapRef readonly markRaw: UnwrapRef readonly nextTick: UnwrapRef + readonly normalizeConfigSource: UnwrapRef readonly onActivated: UnwrapRef readonly onBeforeMount: UnwrapRef readonly onBeforeRouteLeave: UnwrapRef diff --git a/easytier-gui/src/composables/event.ts b/easytier-gui/src/composables/event.ts index 6fb37023..7c8d4e16 100644 --- a/easytier-gui/src/composables/event.ts +++ b/easytier-gui/src/composables/event.ts @@ -14,6 +14,7 @@ const EVENTS = Object.freeze({ PRE_RUN_NETWORK_INSTANCE: 'pre_run_network_instance', POST_RUN_NETWORK_INSTANCE: 'post_run_network_instance', VPN_SERVICE_STOP: 'vpn_service_stop', + VPN_SERVICE_CONFIG_CHANGED: 'vpn_service_config_changed', DHCP_IP_CHANGED: 'dhcp_ip_changed', PROXY_CIDRS_UPDATED: 'proxy_cidrs_updated', EVENT_LAGGED: 'event_lagged', @@ -76,6 +77,14 @@ async function onVpnServiceStop(event: Event) { await syncMobileVpnService(); } +async function onVpnServiceConfigChanged(event: Event) { + const instanceId = normalizeInstanceIdPayload(event.payload) + console.log(`Received event '${EVENTS.VPN_SERVICE_CONFIG_CHANGED}' for instance: ${instanceId}`) + if (type() === 'android') { + await onNetworkInstanceChange(instanceId); + } +} + async function onDhcpIpChanged(event: Event) { const instanceId = normalizeInstanceIdPayload(event.payload) console.log(`Received event '${EVENTS.DHCP_IP_CHANGED}' for instance: ${instanceId}`); @@ -104,6 +113,7 @@ export async function listenGlobalEvents() { await listen(EVENTS.PRE_RUN_NETWORK_INSTANCE, onPreRunNetworkInstance), await listen(EVENTS.POST_RUN_NETWORK_INSTANCE, onPostRunNetworkInstance), await listen(EVENTS.VPN_SERVICE_STOP, onVpnServiceStop), + await listen(EVENTS.VPN_SERVICE_CONFIG_CHANGED, onVpnServiceConfigChanged), await listen(EVENTS.DHCP_IP_CHANGED, onDhcpIpChanged), await listen(EVENTS.PROXY_CIDRS_UPDATED, onProxyCidrsUpdated), await listen(EVENTS.EVENT_LAGGED, onEventLagged), diff --git a/easytier-gui/src/composables/mobile_vpn.ts b/easytier-gui/src/composables/mobile_vpn.ts index a5332608..3d3e2b01 100644 --- a/easytier-gui/src/composables/mobile_vpn.ts +++ b/easytier-gui/src/composables/mobile_vpn.ts @@ -8,17 +8,21 @@ type Route = NetworkTypes.Route interface vpnStatus { running: boolean ipv4Addr: string | null | undefined + ipv4Addrs: string[] ipv4Cidr: number | null | undefined routes: string[] dns: string | null | undefined } let dhcpPollingTimer: NodeJS.Timeout | null = null +let vpnConfigSyncTask: Promise | null = null +let pendingVpnConfigInstanceId: string | null = null const DHCP_POLLING_INTERVAL = 2000 // 2秒后重试 const curVpnStatus: vpnStatus = { running: false, ipv4Addr: undefined, + ipv4Addrs: [], ipv4Cidr: undefined, routes: [], dns: undefined, @@ -42,6 +46,7 @@ async function requestVpnPermission() { function resetVpnConfigStatus() { curVpnStatus.ipv4Addr = undefined + curVpnStatus.ipv4Addrs = [] curVpnStatus.ipv4Cidr = undefined curVpnStatus.routes = [] curVpnStatus.dns = undefined @@ -54,6 +59,12 @@ function syncVpnStatusFromNative(status: Awaited { + onNetworkInstanceChange(instanceId) + }, DHCP_POLLING_INTERVAL) +} + +function hasQueuedVpnConfigChange() { + return pendingVpnConfigInstanceId !== null +} + async function onVpnServiceStart(payload: any) { console.log('vpn service start', JSON.stringify(payload)) curVpnStatus.running = true @@ -170,7 +198,7 @@ function getRoutesForVpn(routes: Route[], node_config: NetworkTypes.NetworkConfi const ret = [] for (const r of routes) { - for (let cidr of r.proxy_cidrs) { + for (let cidr of r.proxy_cidrs ?? []) { if (!cidr.includes('/')) { cidr += '/32' } @@ -178,7 +206,7 @@ function getRoutesForVpn(routes: Route[], node_config: NetworkTypes.NetworkConfi } } - node_config.routes.forEach(r => { + node_config.routes?.forEach(r => { ret.push(r) }) @@ -190,7 +218,32 @@ function getRoutesForVpn(routes: Route[], node_config: NetworkTypes.NetworkConfi return Array.from(new Set(ret)).sort() } +function getCollectedNetworkInfo(response: Awaited>, instanceId: string) { + const info = response.info as any + const map = info?.map ?? info + return map?.[instanceId] +} + export async function onNetworkInstanceChange(instanceId: string) { + pendingVpnConfigInstanceId = instanceId + if (!vpnConfigSyncTask) { + vpnConfigSyncTask = drainVpnConfigChanges().finally(() => { + vpnConfigSyncTask = null + }) + } + + await vpnConfigSyncTask +} + +async function drainVpnConfigChanges() { + while (pendingVpnConfigInstanceId !== null) { + const instanceId = pendingVpnConfigInstanceId + pendingVpnConfigInstanceId = null + await applyNetworkInstanceChange(instanceId) + } +} + +async function applyNetworkInstanceChange(instanceId: string) { console.error('vpn service network instance change id', instanceId) if (dhcpPollingTimer) { @@ -201,60 +254,95 @@ export async function onNetworkInstanceChange(instanceId: string) { if (!instanceId) { console.warn('vpn service skipped because instance id is empty') if (curVpnStatus.running) { + if (hasQueuedVpnConfigChange()) { + return + } await doStopVpn() } return } - const config = await getConfig(instanceId) - console.log('vpn service loaded config', instanceId, JSON.stringify({ - no_tun: config.no_tun, - dhcp: config.dhcp, - enable_magic_dns: config.enable_magic_dns, - })) - if (config.no_tun) { - console.log('vpn service skipped because no_tun is enabled', instanceId) - return - } - const curNetworkInfo = (await collectNetworkInfo(instanceId)).info.map[instanceId] - if (!curNetworkInfo || curNetworkInfo?.error_msg?.length) { - console.warn('vpn service skipped because network info is unavailable', instanceId, curNetworkInfo?.error_msg) + + const group = await findRunningTunInstanceGroup(instanceId) + if (!group.length) { + console.warn('vpn service skipped because no running tun instance is available', instanceId) + if (hasQueuedVpnConfigChange()) { + return + } await doStopVpn() return } - const virtual_ip = Utils.ipv4ToString(curNetworkInfo?.my_node_info?.virtual_ipv4.address) + const ipv4Addrs: string[] = [] + const routes = new Set() + let dns: string | undefined + const retryInstanceIds: string[] = [] + for (const { instanceId, config } of group) { + console.log('vpn service loaded config', instanceId, JSON.stringify({ + no_tun: config.no_tun, + dhcp: config.dhcp, + enable_magic_dns: config.enable_magic_dns, + dev_name: config.dev_name, + })) - if (config.dhcp && (!virtual_ip || !virtual_ip.length)) { - console.log('DHCP enabled but no IP yet, will retry in', DHCP_POLLING_INTERVAL, 'ms') - dhcpPollingTimer = setTimeout(() => { - onNetworkInstanceChange(instanceId) - }, DHCP_POLLING_INTERVAL) - return + const curNetworkInfo = getCollectedNetworkInfo(await collectNetworkInfo(instanceId), instanceId) + if (!curNetworkInfo || curNetworkInfo?.error_msg?.length) { + console.warn('vpn service skipped because network info is unavailable, will retry', instanceId, curNetworkInfo?.error_msg) + retryInstanceIds.push(instanceId) + continue + } + + const virtual_ip = Utils.ipv4ToString(curNetworkInfo?.my_node_info?.virtual_ipv4.address) + if (config.dhcp && (!virtual_ip || !virtual_ip.length)) { + console.log('DHCP enabled but no IP yet, will retry in', DHCP_POLLING_INTERVAL, 'ms') + retryInstanceIds.push(instanceId) + continue + } + + if (!virtual_ip || !virtual_ip.length) { + retryInstanceIds.push(instanceId) + continue + } + + let network_length = curNetworkInfo?.my_node_info?.virtual_ipv4.network_length + if (!network_length) { + network_length = 24 + } + + ipv4Addrs.push(`${virtual_ip}/${network_length}`) + getRoutesForVpn(curNetworkInfo?.routes, config).forEach(route => routes.add(route)) + if (config.enable_magic_dns) { + dns = '100.100.100.101' + } } - if (!virtual_ip || !virtual_ip.length) { + if (retryInstanceIds.length) { + scheduleVpnConfigRetry(retryInstanceIds[0]) + } + + if (!ipv4Addrs.length) { + console.warn('vpn service skipped because no healthy tun instance info is available', instanceId) + if (hasQueuedVpnConfigChange()) { + return + } await doStopVpn() return } - let network_length = curNetworkInfo?.my_node_info?.virtual_ipv4.network_length - if (!network_length) { - network_length = 24 - } - - const routes = getRoutesForVpn(curNetworkInfo?.routes, config) - - const dns = config.enable_magic_dns ? '100.100.100.101' : undefined - - const ipChanged = virtual_ip !== curVpnStatus.ipv4Addr - const cidrChanged = network_length !== curVpnStatus.ipv4Cidr - const routesChanged = JSON.stringify(routes) !== JSON.stringify(curVpnStatus.routes) + const sortedIpv4Addrs = [...ipv4Addrs].sort() + const sortedRoutes = Array.from(routes).sort() + 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 || cidrChanged || routesChanged || dnsChanged + const configChanged = ipChanged || routesChanged || dnsChanged const shouldStartVpn = !curVpnStatus.running if (shouldStartVpn || configChanged) { - console.info('vpn service virtual ip changed', JSON.stringify(curVpnStatus), virtual_ip) + if (hasQueuedVpnConfigChange()) { + console.info('vpn service skipped stale config apply because a newer change is queued') + return + } + + console.info('vpn service virtual ip changed', JSON.stringify(curVpnStatus), sortedIpv4Addrs) if (curVpnStatus.running) { try { await doStopVpn() @@ -262,10 +350,14 @@ export async function onNetworkInstanceChange(instanceId: string) { catch (e) { console.error(e) } + if (hasQueuedVpnConfigChange()) { + console.info('vpn service skipped stale config start because a newer change is queued') + return + } } try { - await doStartVpn(virtual_ip, network_length, routes, dns) + await doStartVpn(sortedIpv4Addrs, sortedRoutes, dns) } catch (e) { if (e instanceof Error && e.message === 'need_prepare') { @@ -304,6 +396,33 @@ async function findRunningTunInstanceId() { return undefined } +async function findRunningTunInstanceGroup(preferredInstanceId?: string) { + const instanceIds = await listNetworkInstanceIds() + const runningIds = instanceIds.running_inst_ids.map(Utils.UuidToStr) + const runningTunInstances = [] + + for (const instanceId of runningIds) { + const config = await getConfig(instanceId) + if (config.no_tun) { + continue + } + runningTunInstances.push({ instanceId, config }) + } + + const selected = runningTunInstances.find(inst => inst.instanceId === preferredInstanceId) + ?? runningTunInstances[0] + if (!selected) { + return [] + } + + const devName = selected.config.dev_name + if (!devName?.length) { + return [selected] + } + + return runningTunInstances.filter(inst => inst.config.dev_name === devName) +} + export async function initMobileVpnService() { await registerVpnServiceListener() } diff --git a/easytier/src/instance/shared_virtual_nic.rs b/easytier/src/instance/shared_virtual_nic.rs index b1766a5e..005cb969 100644 --- a/easytier/src/instance/shared_virtual_nic.rs +++ b/easytier/src/instance/shared_virtual_nic.rs @@ -25,6 +25,7 @@ mod dispatcher; use dispatcher::{SharedVirtualNicDispatcher, SharedVirtualNicMemberTunnelTable}; pub type SharedVirtualNicMemberId = uuid::Uuid; +pub(super) type SharedVirtualNicMemberRegistrationId = uuid::Uuid; #[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] pub struct SharedIpv4Route { @@ -256,6 +257,7 @@ pub struct SharedVirtualNic { ifcfg: SharedIfConfig, valid: Arc, member_tunnel_table: SharedVirtualNicMemberTunnelTable, + member_registrations: BTreeMap, dispatcher: Option, } @@ -266,6 +268,7 @@ impl SharedVirtualNic { ifcfg: SharedIfConfig::default(), valid: Arc::new(AtomicBool::new(true)), member_tunnel_table: SharedVirtualNicMemberTunnelTable::default(), + member_registrations: BTreeMap::new(), dispatcher: None, } } @@ -302,6 +305,71 @@ impl SharedVirtualNic { self.nic.lock().await.link_up().await } + async fn attach_member_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> Result<(), Error> { + self.ensure_valid()?; + + match self.member_registrations.insert(member_id, registration_id) { + Some(old_registration_id) if old_registration_id != registration_id => { + self.remove_member_claims(member_id).await?; + } + _ => {} + } + + Ok(()) + } + + fn is_current_member_registration( + &self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> bool { + self.member_registrations + .get(&member_id) + .is_some_and(|current| *current == registration_id) + } + + async fn apply_member_claims_for_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + self.apply_member_claims(member_id, claims).await + } + + #[cfg(mobile)] + async fn apply_member_claims_for_mobile_registration( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + self.apply_member_claims_for_mobile(member_id, claims).await + } + + async fn remove_member_registration_claims( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> Result<(), Error> { + if !self.is_current_member_registration(member_id, registration_id) { + return Ok(()); + } + + self.member_registrations.remove(&member_id); + self.remove_member_claims(member_id).await + } + async fn apply_member_claims( &mut self, member_id: SharedVirtualNicMemberId, @@ -330,6 +398,25 @@ impl SharedVirtualNic { Ok(()) } + #[cfg(mobile)] + async fn apply_member_claims_for_mobile( + &mut self, + member_id: SharedVirtualNicMemberId, + claims: SharedIfConfigClaims, + ) -> Result<(), Error> { + self.ensure_valid()?; + + let mut next_ifcfg = self.ifcfg.clone(); + let next_claims = claims.clone(); + next_ifcfg.apply_member_claims(member_id, claims); + + self.sync_dispatcher_sources_for_member(member_id, &next_claims) + .await?; + self.ifcfg = next_ifcfg; + + Ok(()) + } + async fn remove_member_claims( &mut self, member_id: SharedVirtualNicMemberId, @@ -340,7 +427,10 @@ impl SharedVirtualNic { let Some(delta) = next_ifcfg.remove_member(member_id) else { return Ok(()); }; + #[cfg(not(mobile))] self.apply_ifcfg_delta(&delta).await?; + #[cfg(mobile)] + drop(delta); self.remove_dispatcher_sources_for_member(member_id).await?; self.ifcfg = next_ifcfg; @@ -430,13 +520,16 @@ impl SharedVirtualNic { ) -> Result<(), Error> { self.ensure_valid()?; - if self.dispatcher.is_some() { + if let Some(dispatcher) = &self.dispatcher { + dispatcher.update_mobile_tun_fd(tun_fd); + self.nic.lock().await.set_mobile_tun_fd_name(tun_fd); return Ok(()); } - let tunnel = self.nic.lock().await.create_dev_for_mobile(tun_fd).await?; - let dispatcher = SharedVirtualNicDispatcher::start( - tunnel, + self.nic.lock().await.set_mobile_tun_fd_name(tun_fd); + let dispatcher = SharedVirtualNicDispatcher::start_for_mobile( + self.nic.clone(), + tun_fd, self.member_tunnel_table.clone(), self.valid.clone(), ); @@ -523,6 +616,7 @@ fn ignore_removed_ifcfg_not_found(result: Result<(), Error>) -> Result<(), Error struct SharedVirtualNicMemberRegistration { member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, shared_nic: Arc>, member_tunnel_table: SharedVirtualNicMemberTunnelTable, } @@ -533,16 +627,22 @@ impl SharedVirtualNicMemberRegistration { tunnel: Box, close_notifier: Arc, ) -> Result<(), Error> { - self.member_tunnel_table - .register(self.member_id, tunnel, close_notifier) + self.member_tunnel_table.register( + self.member_id, + self.registration_id, + tunnel, + close_notifier, + ) } } impl Drop for SharedVirtualNicMemberRegistration { fn drop(&mut self) { - self.member_tunnel_table.unregister(self.member_id); + self.member_tunnel_table + .unregister(self.member_id, self.registration_id); let shared_nic = self.shared_nic.clone(); let member_id = self.member_id; + let registration_id = self.registration_id; let Ok(handle) = tokio::runtime::Handle::try_current() else { tracing::warn!( @@ -554,7 +654,10 @@ impl Drop for SharedVirtualNicMemberRegistration { handle.spawn(async move { let mut shared_nic = shared_nic.lock().await; - if let Err(err) = shared_nic.remove_member_claims(member_id).await { + if let Err(err) = shared_nic + .remove_member_registration_claims(member_id, registration_id) + .await + { tracing::warn!( ?member_id, ?err, @@ -580,12 +683,14 @@ impl SharedVirtualNicMember { close_notifier: Arc, member_tunnel_table: SharedVirtualNicMemberTunnelTable, ) -> Self { + let registration_id = uuid::Uuid::new_v4(); Self { member_id, shared_nic: shared_nic.clone(), close_notifier, registration: Arc::new(SharedVirtualNicMemberRegistration { member_id, + registration_id, shared_nic: shared_nic.clone(), member_tunnel_table, }), @@ -608,6 +713,9 @@ impl SharedVirtualNicMember { let (member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); { let mut shared_nic = self.shared_nic.lock().await; + shared_nic + .attach_member_registration(self.member_id, self.registration.registration_id) + .await?; shared_nic.ensure_dispatcher().await?; } self.registration @@ -623,6 +731,9 @@ impl SharedVirtualNicMember { let (member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); { let mut shared_nic = self.shared_nic.lock().await; + shared_nic + .attach_member_registration(self.member_id, self.registration.registration_id) + .await?; shared_nic.ensure_dispatcher_for_mobile(tun_fd).await?; } self.registration @@ -667,6 +778,24 @@ impl SharedVirtualNicMember { .await } + #[cfg(mobile)] + pub async fn add_mobile_source_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { + let ip = ipv4_inet(ip, cidr)?; + self.update_claims_for_mobile(|claims| { + claims.ipv4_addresses.insert(ip); + }) + .await + } + + #[cfg(mobile)] + pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { + let ip = ipv6_inet(ip, cidr)?; + self.update_claims_for_mobile(|claims| { + claims.ipv6_addresses.insert(ip); + }) + .await + } + pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { self.update_claims(|claims| match ip { Some(ip) => { @@ -740,7 +869,30 @@ impl SharedVirtualNicMember { let mut shared_nic = self.shared_nic.lock().await; let mut claims = shared_nic.ifcfg.claims_of(self.member_id); update(&mut claims); - shared_nic.apply_member_claims(self.member_id, claims).await + shared_nic + .apply_member_claims_for_registration( + self.member_id, + self.registration.registration_id, + claims, + ) + .await + } + + #[cfg(mobile)] + async fn update_claims_for_mobile(&self, update: F) -> Result<(), Error> + where + F: FnOnce(&mut SharedIfConfigClaims) + Send, + { + let mut shared_nic = self.shared_nic.lock().await; + let mut claims = shared_nic.ifcfg.claims_of(self.member_id); + update(&mut claims); + shared_nic + .apply_member_claims_for_mobile_registration( + self.member_id, + self.registration.registration_id, + claims, + ) + .await } } @@ -1081,6 +1233,40 @@ mod tests { drop(shared_nic.nic()); } + #[tokio::test] + async fn stale_member_registration_cleanup_keeps_current_claims() { + let mut shared_nic = SharedVirtualNic::new(virtual_nic_config()); + let member = member_id(1); + let old_registration = uuid::Uuid::from_u128(10); + let current_registration = uuid::Uuid::from_u128(11); + let ip = Ipv4Inet::from_str("10.50.0.2/24").unwrap(); + + shared_nic + .member_registrations + .insert(member, current_registration); + shared_nic.ifcfg_mut().apply_member_claims( + member, + SharedIfConfigClaims { + ipv4_addresses: BTreeSet::from([ip]), + ..Default::default() + }, + ); + + shared_nic + .remove_member_registration_claims(member, old_registration) + .await + .unwrap(); + + assert_eq!( + shared_nic.ifcfg().owners_of_ipv4_address(&ip), + BTreeSet::from([member]) + ); + assert_eq!( + shared_nic.member_registrations.get(&member), + Some(¤t_registration) + ); + } + #[test] fn registry_reuses_shared_virtual_nic_for_same_dev_name() { let mut registry = SharedVirtualNicRegistry::new(); diff --git a/easytier/src/instance/shared_virtual_nic/dispatcher.rs b/easytier/src/instance/shared_virtual_nic/dispatcher.rs index 93bee0f9..b4f7f666 100644 --- a/easytier/src/instance/shared_virtual_nic/dispatcher.rs +++ b/easytier/src/instance/shared_virtual_nic/dispatcher.rs @@ -5,19 +5,28 @@ use std::{ Arc, Mutex as StdMutex, atomic::{AtomicBool, Ordering}, }, + time::Duration, }; use cidr::{Ipv4Inet, Ipv6Inet}; use futures::{SinkExt, StreamExt}; +#[cfg(mobile)] +use std::sync::OnceLock; +#[cfg(mobile)] +use tokio::runtime::{Builder, Runtime}; +#[cfg(mobile)] +use tokio::sync::{Mutex, watch}; use tokio::sync::{Notify, mpsc, oneshot}; use tokio_util::task::AbortOnDropHandle; +#[cfg(mobile)] +use crate::instance::virtual_nic::VirtualNic; use crate::{ common::error::Error, tunnel::{Tunnel, ZCPacketSink, ZCPacketStream, packet_def::ZCPacket}, }; -use super::SharedVirtualNicMemberId; +use super::{SharedVirtualNicMemberId, SharedVirtualNicMemberRegistrationId}; const MEMBER_TUNNEL_BUFFER_SIZE: usize = 1024; const FLOW_OWNER_LIMIT: usize = 4096; @@ -27,6 +36,24 @@ const TCP_HEADER_MIN_LEN: usize = 20; const UDP_HEADER_LEN: usize = 8; const TCP_PROTOCOL: u8 = 6; const UDP_PROTOCOL: u8 = 17; +const DISPATCHER_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(1); +#[cfg(mobile)] +const MOBILE_REBUILD_INITIAL_DELAY: Duration = Duration::from_millis(100); +#[cfg(mobile)] +const MOBILE_REBUILD_MAX_DELAY: Duration = Duration::from_secs(5); + +#[cfg(mobile)] +fn mobile_dispatcher_runtime() -> &'static Runtime { + static RUNTIME: OnceLock = OnceLock::new(); + RUNTIME.get_or_init(|| { + Builder::new_multi_thread() + .worker_threads(1) + .thread_name("easytier-shared-tun") + .enable_all() + .build() + .expect("failed to build shared virtual nic mobile dispatcher runtime") + }) +} struct SharedVirtualNicMemberPacket { member_id: SharedVirtualNicMemberId, @@ -40,12 +67,17 @@ enum SharedVirtualNicControl { }, Unregister { member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, }, UpdateSources { member_id: SharedVirtualNicMemberId, sources: BTreeSet, ack: oneshot::Sender<()>, }, + Shutdown { + invalidate: bool, + ack: oneshot::Sender<()>, + }, } #[derive(Clone, Default)] @@ -60,6 +92,7 @@ struct SharedVirtualNicMemberTunnelTableState { } struct SharedVirtualNicMemberTunnelEntry { + registration_id: SharedVirtualNicMemberRegistrationId, sender: mpsc::Sender, close_notifier: Arc, _tasks: Vec>, @@ -76,7 +109,7 @@ impl SharedVirtualNicMemberTunnelTable { state.control_sender = Some(control_sender); } - fn detach_dispatcher(&self) { + pub(super) fn detach_dispatcher(&self) { let mut state = self.state.lock().unwrap(); state.to_tun_sender.take(); state.control_sender.take(); @@ -85,6 +118,7 @@ impl SharedVirtualNicMemberTunnelTable { pub(super) fn register( &self, member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, tunnel: Box, close_notifier: Arc, ) -> Result<(), Error> { @@ -121,7 +155,10 @@ impl SharedVirtualNicMemberTunnelTable { } } - let _ = reader_control_sender.send(SharedVirtualNicControl::Unregister { member_id }); + let _ = reader_control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); reader_close_notifier.notify_one(); })); @@ -131,8 +168,10 @@ impl SharedVirtualNicMemberTunnelTable { while let Some(packet) = to_member_receiver.recv().await { if let Err(err) = member_sink.send(packet).await { tracing::error!(?member_id, ?err, "shared member tunnel write failed"); - let _ = writer_control_sender - .send(SharedVirtualNicControl::Unregister { member_id }); + let _ = writer_control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); writer_close_notifier.notify_one(); break; } @@ -140,6 +179,7 @@ impl SharedVirtualNicMemberTunnelTable { })); let entry = SharedVirtualNicMemberTunnelEntry { + registration_id, sender: to_member_sender, close_notifier, _tasks: vec![reader_task, writer_task], @@ -152,11 +192,18 @@ impl SharedVirtualNicMemberTunnelTable { Ok(()) } - pub(super) fn unregister(&self, member_id: SharedVirtualNicMemberId) { + pub(super) fn unregister( + &self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) { let Some(control_sender) = self.control_sender() else { return; }; - let _ = control_sender.send(SharedVirtualNicControl::Unregister { member_id }); + let _ = control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); } fn dispatcher_channels( @@ -309,6 +356,8 @@ impl SharedVirtualNicFlowTable { pub(super) struct SharedVirtualNicDispatcher { _task: AbortOnDropHandle<()>, control_sender: mpsc::UnboundedSender, + #[cfg(mobile)] + mobile_tun_fd_sender: Option>, } impl SharedVirtualNicDispatcher { @@ -335,6 +384,51 @@ impl SharedVirtualNicDispatcher { Self { _task: AbortOnDropHandle::new(tokio::spawn(task.run())), control_sender, + #[cfg(mobile)] + mobile_tun_fd_sender: None, + } + } + + #[cfg(mobile)] + pub(super) fn start_for_mobile( + nic: Arc>, + tun_fd: std::os::fd::RawFd, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + ) -> Self { + let (to_tun_sender, to_tun_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (control_sender, control_receiver) = mpsc::unbounded_channel(); + let (mobile_tun_fd_sender, mobile_tun_fd_receiver) = watch::channel(tun_fd); + let task_member_tunnel_table = member_tunnel_table.clone(); + member_tunnel_table.attach_dispatcher(to_tun_sender, control_sender.clone()); + + let task = mobile_dispatcher_runtime().spawn(async move { + let task = SharedVirtualNicMobileDispatcherTask { + nic, + tun_stream: None, + tun_sink: None, + tun_fd: mobile_tun_fd_receiver, + to_tun_receiver, + control_receiver, + member_tunnel_table: task_member_tunnel_table, + valid, + state: SharedVirtualNicDispatcherState::default(), + }; + + task.run().await; + }); + + Self { + _task: AbortOnDropHandle::new(task), + control_sender, + mobile_tun_fd_sender: Some(mobile_tun_fd_sender), + } + } + + #[cfg(mobile)] + pub(super) fn update_mobile_tun_fd(&self, tun_fd: std::os::fd::RawFd) { + if let Some(sender) = &self.mobile_tun_fd_sender { + sender.send_replace(tun_fd); } } @@ -371,6 +465,27 @@ impl SharedVirtualNicDispatcher { rx.await .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running").into()) } + + pub(super) async fn shutdown_without_invalidation(self) { + let (ack, rx) = oneshot::channel(); + if self + .control_sender + .send(SharedVirtualNicControl::Shutdown { + invalidate: false, + ack, + }) + .is_err() + { + return; + } + + if tokio::time::timeout(DISPATCHER_SHUTDOWN_TIMEOUT, rx) + .await + .is_err() + { + tracing::warn!("timed out shutting down shared virtual nic dispatcher"); + } + } } struct SharedVirtualNicDispatcherTask { @@ -391,7 +506,14 @@ impl SharedVirtualNicDispatcherTask { let Some(control) = control else { break; }; - self.state.handle_control(control); + match control { + SharedVirtualNicControl::Shutdown { invalidate, ack } => { + self.cleanup(invalidate); + let _ = ack.send(()); + return; + } + other => self.state.handle_control(other), + } } member_packet = self.to_tun_receiver.recv() => { let Some(member_packet) = member_packet else { @@ -417,7 +539,13 @@ impl SharedVirtualNicDispatcherTask { } } - self.valid.store(false, Ordering::Release); + self.cleanup(true); + } + + fn cleanup(&mut self, invalidate: bool) { + if invalidate { + self.valid.store(false, Ordering::Release); + } self.member_tunnel_table.detach_dispatcher(); self.state.close_all(); } @@ -436,6 +564,207 @@ impl SharedVirtualNicDispatcherTask { } } +#[cfg(mobile)] +struct SharedVirtualNicMobileDispatcherTask { + nic: Arc>, + tun_stream: Option>>, + tun_sink: Option>>, + tun_fd: watch::Receiver, + to_tun_receiver: mpsc::Receiver, + control_receiver: mpsc::UnboundedReceiver, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + state: SharedVirtualNicDispatcherState, +} + +#[cfg(mobile)] +impl SharedVirtualNicMobileDispatcherTask { + async fn run(mut self) { + let mut rebuild_delay = MOBILE_REBUILD_INITIAL_DELAY; + let mut wait_before_rebuild = false; + let mut rebuild_deadline = None; + + loop { + if self.tun_stream.is_none() { + if !wait_before_rebuild { + if !self.rebuild_tun().await { + wait_before_rebuild = true; + rebuild_deadline = None; + } + continue; + } + + let deadline = *rebuild_deadline + .get_or_insert_with(|| tokio::time::Instant::now() + rebuild_delay); + tokio::select! { + control = self.control_receiver.recv() => { + if !self.handle_control(control) { + return; + } + } + member_packet = self.to_tun_receiver.recv() => { + let Some(member_packet) = member_packet else { + self.cleanup(true); + return; + }; + tracing::trace!( + member_id = ?member_packet.member_id, + "shared virtual nic dropped member packet while rebuilding mobile tun" + ); + } + changed = self.tun_fd.changed() => { + if changed.is_err() { + self.cleanup(true); + return; + } + rebuild_delay = MOBILE_REBUILD_INITIAL_DELAY; + wait_before_rebuild = false; + rebuild_deadline = None; + } + _ = tokio::time::sleep_until(deadline) => { + rebuild_delay = next_mobile_rebuild_delay(rebuild_delay); + wait_before_rebuild = false; + rebuild_deadline = None; + } + } + continue; + } + + tokio::select! { + control = self.control_receiver.recv() => { + if !self.handle_control(control) { + return; + } + } + member_packet = self.to_tun_receiver.recv() => { + let Some(member_packet) = member_packet else { + self.cleanup(true); + return; + }; + if self.forward_member_packet_to_tun(member_packet).await { + wait_before_rebuild = true; + rebuild_deadline = None; + } else { + rebuild_delay = MOBILE_REBUILD_INITIAL_DELAY; + } + } + packet = self.tun_stream.as_mut().expect("mobile tun stream should exist").next() => { + let Some(packet) = packet else { + tracing::error!("shared virtual nic mobile tun stream closed"); + self.drop_tun(); + wait_before_rebuild = true; + rebuild_deadline = None; + continue; + }; + let packet = match packet { + Ok(packet) => packet, + Err(err) => { + tracing::error!(?err, "shared virtual nic read from mobile tun failed"); + self.drop_tun(); + wait_before_rebuild = true; + rebuild_deadline = None; + continue; + } + }; + rebuild_delay = MOBILE_REBUILD_INITIAL_DELAY; + self.state.forward_tun_packet_to_member(packet).await; + } + changed = self.tun_fd.changed() => { + if changed.is_err() { + self.cleanup(true); + return; + } + self.drop_tun(); + rebuild_delay = MOBILE_REBUILD_INITIAL_DELAY; + wait_before_rebuild = false; + rebuild_deadline = None; + } + } + } + } + + async fn rebuild_tun(&mut self) -> bool { + let tun_fd = *self.tun_fd.borrow_and_update(); + match self.nic.lock().await.create_dev_for_mobile(tun_fd).await { + Ok(tunnel) => { + let (tun_stream, tun_sink) = tunnel.split(); + self.tun_stream = Some(tun_stream); + self.tun_sink = Some(tun_sink); + tracing::info!(fd = tun_fd, "rebuilt shared virtual nic mobile tun"); + true + } + Err(err) => { + tracing::error!( + fd = tun_fd, + ?err, + "failed to rebuild shared virtual nic mobile tun" + ); + false + } + } + } + + async fn forward_member_packet_to_tun( + &mut self, + member_packet: SharedVirtualNicMemberPacket, + ) -> bool { + self.state + .remember_reverse_owner(member_packet.member_id, &member_packet.packet); + let Some(tun_sink) = self.tun_sink.as_mut() else { + tracing::trace!( + member_id = ?member_packet.member_id, + "shared virtual nic dropped member packet without mobile tun" + ); + return false; + }; + + if let Err(err) = tun_sink.send(member_packet.packet).await { + tracing::error!(?err, "shared virtual nic write to mobile tun failed"); + self.drop_tun(); + return true; + } + false + } + + fn handle_control(&mut self, control: Option) -> bool { + let Some(control) = control else { + self.cleanup(true); + return false; + }; + + match control { + SharedVirtualNicControl::Shutdown { invalidate, ack } => { + self.cleanup(invalidate); + let _ = ack.send(()); + false + } + other => { + self.state.handle_control(other); + true + } + } + } + + fn drop_tun(&mut self) { + self.tun_stream.take(); + self.tun_sink.take(); + } + + fn cleanup(&mut self, invalidate: bool) { + self.drop_tun(); + if invalidate { + self.valid.store(false, Ordering::Release); + } + self.member_tunnel_table.detach_dispatcher(); + self.state.close_all(); + } +} + +#[cfg(mobile)] +fn next_mobile_rebuild_delay(delay: Duration) -> Duration { + delay.saturating_mul(2).min(MOBILE_REBUILD_MAX_DELAY) +} + #[derive(Default)] struct SharedVirtualNicDispatcherState { members: BTreeMap, @@ -449,8 +778,11 @@ impl SharedVirtualNicDispatcherState { SharedVirtualNicControl::Register { member_id, entry } => { self.register(member_id, entry); } - SharedVirtualNicControl::Unregister { member_id } => { - self.unregister(member_id); + SharedVirtualNicControl::Unregister { + member_id, + registration_id, + } => { + self.unregister(member_id, registration_id); } SharedVirtualNicControl::UpdateSources { member_id, @@ -460,6 +792,9 @@ impl SharedVirtualNicDispatcherState { self.source_table.update_member_sources(member_id, sources); let _ = ack.send(()); } + SharedVirtualNicControl::Shutdown { .. } => { + unreachable!("dispatcher shutdown is handled by the dispatcher task") + } } } @@ -472,7 +807,20 @@ impl SharedVirtualNicDispatcherState { drop(old_entry); } - fn unregister(&mut self, member_id: SharedVirtualNicMemberId) { + fn unregister( + &mut self, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, + ) { + if self + .members + .get(&member_id) + .map(|entry| entry.registration_id) + != Some(registration_id) + { + return; + } + let entry = self.members.remove(&member_id); drop(entry); self.flow_table.remove_owner(member_id); @@ -530,10 +878,10 @@ impl SharedVirtualNicDispatcherState { member_id: SharedVirtualNicMemberId, packet: ZCPacket, ) -> Result<(), ZCPacket> { - let Some(sender) = self + let Some((registration_id, sender)) = self .members .get(&member_id) - .map(|entry| entry.sender.clone()) + .map(|entry| (entry.registration_id, entry.sender.clone())) else { return Err(packet); }; @@ -541,7 +889,7 @@ impl SharedVirtualNicDispatcherState { match sender.send(packet).await { Ok(()) => Ok(()), Err(err) => { - self.unregister(member_id); + self.unregister(member_id, registration_id); Err(err.0) } } @@ -690,7 +1038,15 @@ mod tests { } fn member_entry(sender: mpsc::Sender) -> SharedVirtualNicMemberTunnelEntry { + member_entry_with_registration(sender, uuid::Uuid::from_u128(1)) + } + + fn member_entry_with_registration( + sender: mpsc::Sender, + registration_id: SharedVirtualNicMemberRegistrationId, + ) -> SharedVirtualNicMemberTunnelEntry { SharedVirtualNicMemberTunnelEntry { + registration_id, sender, close_notifier: Arc::new(Notify::new()), _tasks: Vec::new(), @@ -772,7 +1128,7 @@ mod tests { owner, BTreeSet::from([SharedVirtualNicFlowAddr::V6(source.octets())]), ); - state.unregister(owner); + state.unregister(owner, uuid::Uuid::from_u128(1)); state .forward_tun_packet_to_member(ipv6_packet(source, dst)) .await; @@ -780,6 +1136,28 @@ mod tests { assert!(fallback_receiver.try_recv().is_err()); } + #[tokio::test] + async fn dispatcher_ignores_stale_member_unregister() { + let member_id = uuid::Uuid::from_u128(1); + let stale_registration = uuid::Uuid::from_u128(10); + let current_registration = uuid::Uuid::from_u128(11); + let src = "2001:db8::1".parse::().unwrap(); + let dst = "2001:db8:ffff::1".parse::().unwrap(); + let (sender, mut receiver) = mpsc::channel(1); + let mut state = SharedVirtualNicDispatcherState::default(); + + state.register( + member_id, + member_entry_with_registration(sender, current_registration), + ); + state.unregister(member_id, stale_registration); + state + .forward_tun_packet_to_member(ipv6_packet(src, dst)) + .await; + + assert!(receiver.try_recv().is_ok()); + } + #[tokio::test] async fn dispatcher_invalidates_shared_nic_when_tun_read_fails() { let member_id = uuid::Uuid::from_u128(1); @@ -800,7 +1178,12 @@ mod tests { let (_member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); member_tunnel_table - .register(member_id, shared_tunnel, close_notifier.clone()) + .register( + member_id, + uuid::Uuid::from_u128(1), + shared_tunnel, + close_notifier.clone(), + ) .unwrap(); dispatcher .update_sources(member_id, &BTreeSet::new(), &BTreeSet::new()) @@ -815,4 +1198,45 @@ mod tests { assert!(!valid.load(Ordering::Acquire)); assert!(member_tunnel_table.dispatcher_channels().is_none()); } + + #[tokio::test] + async fn dispatcher_shutdown_for_replacement_keeps_shared_nic_valid() { + let member_id = uuid::Uuid::from_u128(1); + let (_tun_tx, tun_rx) = mpsc::unbounded_channel(); + let tun_stream = tokio_stream::wrappers::UnboundedReceiverStream::new(tun_rx); + let tun_sink = futures::sink::unfold((), |(), _packet: ZCPacket| async { + Ok::<(), TunnelError>(()) + }); + let tunnel = TunnelWrapper::new(tun_stream, tun_sink, None); + let member_tunnel_table = SharedVirtualNicMemberTunnelTable::default(); + let valid = Arc::new(AtomicBool::new(true)); + let dispatcher = SharedVirtualNicDispatcher::start( + Box::new(tunnel), + member_tunnel_table.clone(), + valid.clone(), + ); + let close_notifier = Arc::new(Notify::new()); + let (_member_tunnel, shared_tunnel) = create_ring_tunnel_pair(); + + member_tunnel_table + .register( + member_id, + uuid::Uuid::from_u128(1), + shared_tunnel, + close_notifier.clone(), + ) + .unwrap(); + dispatcher + .update_sources(member_id, &BTreeSet::new(), &BTreeSet::new()) + .await + .unwrap(); + + dispatcher.shutdown_without_invalidation().await; + + tokio::time::timeout(Duration::from_secs(1), close_notifier.notified()) + .await + .unwrap(); + assert!(valid.load(Ordering::Acquire)); + assert!(member_tunnel_table.dispatcher_channels().is_none()); + } } diff --git a/easytier/src/instance/virtual_nic.rs b/easytier/src/instance/virtual_nic.rs index 0ee7acef..e827599d 100644 --- a/easytier/src/instance/virtual_nic.rs +++ b/easytier/src/instance/virtual_nic.rs @@ -599,6 +599,11 @@ impl VirtualNic { Ok(tun::create(&config)?) } + #[cfg(mobile)] + pub fn set_mobile_tun_fd_name(&mut self, tun_fd: std::os::fd::RawFd) { + self.ifname = Some(format!("tunfd_{}", tun_fd)); + } + #[cfg(mobile)] pub async fn create_dev_for_mobile( &mut self, @@ -634,7 +639,7 @@ impl VirtualNic { None, ); - self.ifname = Some(format!("tunfd_{}", tun_fd)); + self.set_mobile_tun_fd_name(tun_fd); Ok(Box::new(ft)) } @@ -986,6 +991,22 @@ impl NicBackend { Self::Shared(member) => member.add_ipv6(ip, cidr).await, } } + + #[cfg(mobile)] + pub async fn add_mobile_source_ip(&self, ip: Ipv4Addr, cidr: i32) -> Result<(), Error> { + match self { + Self::Dedicated(_) => Ok(()), + Self::Shared(member) => member.add_mobile_source_ip(ip, cidr).await, + } + } + + #[cfg(mobile)] + pub async fn add_mobile_source_ipv6(&self, ip: Ipv6Addr, cidr: i32) -> Result<(), Error> { + match self { + Self::Dedicated(_) => Ok(()), + Self::Shared(member) => member.add_mobile_source_ipv6(ip, cidr).await, + } + } } pub struct NicCtx { @@ -1673,16 +1694,14 @@ impl NicCtx { #[cfg(mobile)] pub async fn run_for_mobile(&mut self, tun_fd: std::os::fd::RawFd) -> Result<(), Error> { - let tunnel = match self.backend.create_dev_for_mobile(tun_fd).await { + let (tunnel, ifname) = match self.backend.create_dev_for_mobile(tun_fd).await { Ok(ret) => { let ifname = self .backend .ifname() .await .ok_or_else(|| anyhow::anyhow!("tun device has no interface name"))?; - self.global_ctx - .issue_event(GlobalCtxEvent::TunDeviceReady(ifname)); - ret + (ret, ifname) } Err(err) => { self.global_ctx @@ -1691,6 +1710,20 @@ impl NicCtx { } }; + if let Some(ipv4_addr) = self.global_ctx.get_ipv4() { + self.backend + .add_mobile_source_ip(ipv4_addr.address(), ipv4_addr.network_length() as i32) + .await?; + } + if let Some(ipv6_addr) = self.global_ctx.get_ipv6() { + self.backend + .add_mobile_source_ipv6(ipv6_addr.address(), ipv6_addr.network_length() as i32) + .await?; + } + + self.global_ctx + .issue_event(GlobalCtxEvent::TunDeviceReady(ifname)); + let (stream, sink) = tunnel.split(); self.do_forward_nic_to_peers_task(stream)?; diff --git a/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt b/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt index b1827da0..41b77e5f 100644 --- a/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt +++ b/tauri-plugin-vpnservice/android/src/main/java/TauriVpnService.kt @@ -15,10 +15,12 @@ class TauriVpnService : VpnService() { @JvmField var triggerCallback: (String, JSObject) -> Unit = { _, _ -> } @JvmField var self: TauriVpnService? = null @JvmField var ipv4Addr: String? = null + @JvmField var ipv4Addrs: Array = emptyArray() @JvmField var routes: Array = emptyArray() @JvmField var dns: String? = null const val IPV4_ADDR = "IPV4_ADDR" + const val IPV4_ADDRS = "IPV4_ADDRS" const val ROUTES = "ROUTES" const val DNS = "DNS" const val DISALLOWED_APPLICATIONS = "DISALLOWED_APPLICATIONS" @@ -30,7 +32,8 @@ class TauriVpnService : VpnService() { override fun onStartCommand(intent: Intent?, flags: Int, startId: Int): Int { println("vpn on start command ${intent?.getExtras()} $intent") var args = intent?.getExtras() - ipv4Addr = args?.getString(IPV4_ADDR) + ipv4Addrs = getIpv4Addrs(args) + ipv4Addr = ipv4Addrs.firstOrNull() routes = args?.getStringArray(ROUTES) ?: emptyArray() dns = args?.getString(DNS) @@ -74,28 +77,44 @@ class TauriVpnService : VpnService() { private fun clearStatus() { ipv4Addr = null + ipv4Addrs = emptyArray() routes = emptyArray() dns = null } + private fun getIpv4Addrs(args: Bundle?): Array { + val ipv4Addrs = args + ?.getStringArray(IPV4_ADDRS) + ?.filter { it.isNotBlank() } + ?.toTypedArray() + ?: emptyArray() + if (ipv4Addrs.isNotEmpty()) { + return ipv4Addrs + } + + return arrayOf(args?.getString(IPV4_ADDR) ?: "10.126.126.1/24") + } + private fun createVpnInterface(args: Bundle?): ParcelFileDescriptor { var builder = Builder() .setSession("TauriVpnService") .setBlocking(false) var mtu = args?.getInt(MTU) ?: 1500 - var ipv4Addr = args?.getString(IPV4_ADDR) ?: "10.126.126.1/24" + var ipv4Addrs = getIpv4Addrs(args) var dns: String? = args?.getString(DNS) var routes = args?.getStringArray(ROUTES) ?: emptyArray() var disallowedApplications = args?.getStringArray(DISALLOWED_APPLICATIONS) ?: emptyArray() - println("vpn create vpn interface. mtu: $mtu, ipv4Addr: $ipv4Addr, dns:" + + println("vpn create vpn interface. mtu: $mtu, ipv4Addrs: ${java.util.Arrays.toString(ipv4Addrs)}, dns:" + "$dns, routes: ${java.util.Arrays.toString(routes)}," + "disallowedApplications: ${java.util.Arrays.toString(disallowedApplications)}") - val ipParts = ipv4Addr.split("/") - if (ipParts.size != 2) throw IllegalArgumentException("Invalid IP addr string") - builder.addAddress(ipParts[0], ipParts[1].toInt()) + for (ipv4Addr in ipv4Addrs) { + val ipParts = ipv4Addr.split("/") + if (ipParts.size != 2) throw IllegalArgumentException("Invalid IP addr string") + builder.addAddress(ipParts[0], ipParts[1].toInt()) + } builder.addAddress("fd00::1", 128) builder.setMtu(mtu) diff --git a/tauri-plugin-vpnservice/android/src/main/java/VpnServicePlugin.kt b/tauri-plugin-vpnservice/android/src/main/java/VpnServicePlugin.kt index abd4a23f..673c980a 100644 --- a/tauri-plugin-vpnservice/android/src/main/java/VpnServicePlugin.kt +++ b/tauri-plugin-vpnservice/android/src/main/java/VpnServicePlugin.kt @@ -12,6 +12,7 @@ import app.tauri.plugin.Invoke import app.tauri.plugin.JSObject import app.tauri.plugin.Plugin import android.webkit.WebView +import org.json.JSONArray @InvokeArg class PingArgs { @@ -21,6 +22,7 @@ class PingArgs { @InvokeArg class StartVpnArgs { var ipv4Addr: String? = null + var ipv4Addrs: Array = emptyArray() var routes: Array = emptyArray() var dns: String? = null var disallowedApplications: Array = emptyArray() @@ -85,6 +87,7 @@ class VpnServicePlugin(private val activity: Activity) : Plugin(activity) { } else { val intent = Intent(activity, TauriVpnService::class.java) intent.putExtra(TauriVpnService.IPV4_ADDR, args.ipv4Addr) + intent.putExtra(TauriVpnService.IPV4_ADDRS, args.ipv4Addrs) intent.putExtra(TauriVpnService.ROUTES, args.routes) intent.putExtra(TauriVpnService.DNS, args.dns) intent.putExtra(TauriVpnService.DISALLOWED_APPLICATIONS, args.disallowedApplications) @@ -112,7 +115,8 @@ class VpnServicePlugin(private val activity: Activity) : Plugin(activity) { val ret = JSObject() ret.put("running", TauriVpnService.self != null) ret.put("ipv4Addr", TauriVpnService.ipv4Addr) - ret.put("routes", TauriVpnService.routes) + ret.put("ipv4Addrs", JSONArray(TauriVpnService.ipv4Addrs)) + ret.put("routes", JSONArray(TauriVpnService.routes)) ret.put("dns", TauriVpnService.dns) invoke.resolve(ret) } diff --git a/tauri-plugin-vpnservice/guest-js/index.ts b/tauri-plugin-vpnservice/guest-js/index.ts index 1ce6d662..16d07a9b 100644 --- a/tauri-plugin-vpnservice/guest-js/index.ts +++ b/tauri-plugin-vpnservice/guest-js/index.ts @@ -15,6 +15,7 @@ export interface InvokeResponse { export interface StartVpnRequest { ipv4Addr?: string; + ipv4Addrs?: string[]; routes?: string[]; dns?: string; disallowedApplications?: string[]; @@ -24,6 +25,7 @@ export interface StartVpnRequest { export interface VpnStatusResponse { running: boolean; ipv4Addr?: string; + ipv4Addrs?: string[]; routes?: string[]; dns?: string; } diff --git a/tauri-plugin-vpnservice/src/models.rs b/tauri-plugin-vpnservice/src/models.rs index 7c1716d8..9fde47e1 100644 --- a/tauri-plugin-vpnservice/src/models.rs +++ b/tauri-plugin-vpnservice/src/models.rs @@ -22,6 +22,7 @@ pub struct VoidRequest {} #[serde(rename_all = "camelCase")] pub struct StartVpnRequest { pub ipv4_addr: Option, + pub ipv4_addrs: Option>, pub routes: Option>, pub dns: Option, pub disallowed_applications: Option>, @@ -39,6 +40,7 @@ pub struct Status { pub struct VpnStatus { pub running: bool, pub ipv4_addr: Option, + pub ipv4_addrs: Option>, pub routes: Option>, pub dns: Option, }