diff --git a/easytier/src/instance/instance.rs b/easytier/src/instance/instance.rs index c0cbbf74..96065242 100644 --- a/easytier/src/instance/instance.rs +++ b/easytier/src/instance/instance.rs @@ -904,6 +904,21 @@ impl Instance { tracing::debug!("nic ctx updated."); } + #[cfg(all(feature = "tun", feature = "magic-dns"))] + async fn use_new_nic_ctx_with_magic_dns( + arc_nic_ctx: ArcNicCtx, + peer_mgr: Arc, + nic_ctx: NicCtx, + tun_dev: Option, + tun_ip: Option, + ) { + let route_backend = nic_ctx.shared_route_backend_for_dns(); + let magic_dns = tun_ip.and_then(|tun_ip| { + Self::create_magic_dns_runner(peer_mgr, tun_dev, tun_ip, route_backend) + }); + Self::use_new_nic_ctx(arc_nic_ctx, nic_ctx, magic_dns).await; + } + #[cfg(feature = "tun")] async fn new_nic_ctx( global_ctx: ArcGlobalCtx, @@ -1052,21 +1067,19 @@ impl Instance { continue; } #[cfg(feature = "magic-dns")] - let route_backend = new_nic_ctx.shared_route_backend_for_dns(); - #[cfg(feature = "magic-dns")] - let ifname = new_nic_ctx.ifname().await; - Self::use_new_nic_ctx( - nic_ctx.clone(), - new_nic_ctx, - #[cfg(feature = "magic-dns")] - Self::create_magic_dns_runner( + { + let ifname = new_nic_ctx.ifname().await; + Self::use_new_nic_ctx_with_magic_dns( + nic_ctx.clone(), peer_manager_c.clone(), + new_nic_ctx, ifname, - ip, - route_backend, - ), - ) - .await; + Some(ip), + ) + .await; + } + #[cfg(not(feature = "magic-dns"))] + Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx).await; } current_dhcp_ip = Some(ip); @@ -1145,14 +1158,15 @@ impl Instance { // Create Magic DNS runner only if we have IPv4 #[cfg(feature = "magic-dns")] { - let route_backend = new_nic_ctx.shared_route_backend_for_dns(); let ifname = new_nic_ctx.ifname().await; - let dns_runner = if let Some(ipv4) = ipv4_addr { - Self::create_magic_dns_runner(peer_mgr, ifname, ipv4, route_backend) - } else { - None - }; - Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx, dns_runner).await; + Self::use_new_nic_ctx_with_magic_dns( + nic_ctx.clone(), + peer_mgr, + new_nic_ctx, + ifname, + ipv4_addr, + ) + .await; } #[cfg(not(feature = "magic-dns"))] Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx).await; @@ -1757,13 +1771,14 @@ impl Instance { #[cfg(feature = "magic-dns")] { - let route_backend = new_nic_ctx.shared_route_backend_for_dns(); - let magic_dns_runner = if let Some(ipv4) = global_ctx.get_ipv4() { - Self::create_magic_dns_runner(peer_manager.clone(), None, ipv4, route_backend) - } else { - None - }; - Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx, magic_dns_runner).await; + Self::use_new_nic_ctx_with_magic_dns( + nic_ctx.clone(), + peer_manager.clone(), + new_nic_ctx, + None, + global_ctx.get_ipv4(), + ) + .await; } #[cfg(not(feature = "magic-dns"))] Self::use_new_nic_ctx(nic_ctx.clone(), new_nic_ctx).await; diff --git a/easytier/src/instance/shared_virtual_nic.rs b/easytier/src/instance/shared_virtual_nic.rs index 07c1aa4e..ecc09ef3 100644 --- a/easytier/src/instance/shared_virtual_nic.rs +++ b/easytier/src/instance/shared_virtual_nic.rs @@ -103,6 +103,13 @@ pub struct SharedIfConfigDelta { pub mtu: Option, } +struct SharedIfConfigOwnerDeltas { + ipv4_addresses: OwnedItemDelta, + ipv6_addresses: OwnedItemDelta, + ipv4_routes: OwnedItemDelta, + ipv6_routes: OwnedItemDelta, +} + #[derive(Clone, Debug, Default, PartialEq, Eq)] pub struct SharedIfConfigSnapshot { pub ipv4_addresses: BTreeMap>, @@ -134,59 +141,13 @@ 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 source_change_candidates = merged_items(&old_claims.ipv4_routes, &claims.ipv4_routes); let old_ipv4_route_sources = self.ipv4_route_sources(&source_change_candidates); - - let ipv4_addresses = update_owned_items( - &mut self.ipv4_address_owners, - member_id, - &old_claims.ipv4_addresses, - &claims.ipv4_addresses, - ); - let ipv6_addresses = update_owned_items( - &mut self.ipv6_address_owners, - member_id, - &old_claims.ipv6_addresses, - &claims.ipv6_addresses, - ); - let ipv4_routes = update_owned_items( - &mut self.ipv4_route_owners, - member_id, - &old_claims.ipv4_routes, - &claims.ipv4_routes, - ); - let ipv6_routes = update_owned_items( - &mut self.ipv6_route_owners, - member_id, - &old_claims.ipv6_routes, - &claims.ipv6_routes, - ); + let owner_deltas = self.update_owner_deltas(member_id, &old_claims, &claims); 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()), - } + self.build_delta(old_mtu, &old_ipv4_route_sources, owner_deltas) } pub fn remove_member( @@ -197,48 +158,100 @@ impl SharedIfConfig { 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, - member_id, - &old_claims.ipv4_addresses, - ); - let ipv6_addresses = remove_owned_items( - &mut self.ipv6_address_owners, - member_id, - &old_claims.ipv6_addresses, - ); - let ipv4_routes = remove_owned_items( - &mut self.ipv4_route_owners, - member_id, - &old_claims.ipv4_routes, - ); - let ipv6_routes = remove_owned_items( - &mut self.ipv6_route_owners, - member_id, - &old_claims.ipv6_routes, - ); - + let owner_deltas = self.remove_owner_deltas(member_id, &old_claims); self.member_mtu.remove(&member_id); + + Some(self.build_delta(old_mtu, &old_ipv4_route_sources, owner_deltas)) + } + + fn update_owner_deltas( + &mut self, + member_id: SharedVirtualNicMemberId, + old_claims: &SharedIfConfigClaims, + claims: &SharedIfConfigClaims, + ) -> SharedIfConfigOwnerDeltas { + SharedIfConfigOwnerDeltas { + ipv4_addresses: update_owned_items( + &mut self.ipv4_address_owners, + member_id, + &old_claims.ipv4_addresses, + &claims.ipv4_addresses, + ), + ipv6_addresses: update_owned_items( + &mut self.ipv6_address_owners, + member_id, + &old_claims.ipv6_addresses, + &claims.ipv6_addresses, + ), + ipv4_routes: update_owned_items( + &mut self.ipv4_route_owners, + member_id, + &old_claims.ipv4_routes, + &claims.ipv4_routes, + ), + ipv6_routes: update_owned_items( + &mut self.ipv6_route_owners, + member_id, + &old_claims.ipv6_routes, + &claims.ipv6_routes, + ), + } + } + + fn remove_owner_deltas( + &mut self, + member_id: SharedVirtualNicMemberId, + old_claims: &SharedIfConfigClaims, + ) -> SharedIfConfigOwnerDeltas { + SharedIfConfigOwnerDeltas { + ipv4_addresses: remove_owned_items( + &mut self.ipv4_address_owners, + member_id, + &old_claims.ipv4_addresses, + ), + ipv6_addresses: remove_owned_items( + &mut self.ipv6_address_owners, + member_id, + &old_claims.ipv6_addresses, + ), + ipv4_routes: remove_owned_items( + &mut self.ipv4_route_owners, + member_id, + &old_claims.ipv4_routes, + ), + ipv6_routes: remove_owned_items( + &mut self.ipv6_route_owners, + member_id, + &old_claims.ipv6_routes, + ), + } + } + + fn build_delta( + &self, + old_mtu: Option, + old_ipv4_route_sources: &BTreeMap>, + owner_deltas: SharedIfConfigOwnerDeltas, + ) -> SharedIfConfigDelta { let ipv4_route_removed_old_source_hints = - old_ipv4_route_hints(&old_ipv4_route_sources, &ipv4_routes.removed); + old_ipv4_route_hints(old_ipv4_route_sources, &owner_deltas.ipv4_routes.removed); let ipv4_route_source_changed_old_hints = - self.changed_ipv4_route_sources(&old_ipv4_route_sources, &ipv4_routes); + self.changed_ipv4_route_sources(old_ipv4_route_sources, &owner_deltas.ipv4_routes); let ipv4_route_source_changed = ipv4_route_source_changed_old_hints .keys() .cloned() .collect(); - Some(SharedIfConfigDelta { - ipv4_addresses, - ipv6_addresses, - ipv4_routes, + SharedIfConfigDelta { + ipv4_addresses: owner_deltas.ipv4_addresses, + ipv6_addresses: owner_deltas.ipv6_addresses, + ipv4_routes: owner_deltas.ipv4_routes, ipv4_route_removed_old_source_hints, ipv4_route_source_changed, ipv4_route_source_changed_old_hints, - ipv6_routes, + ipv6_routes: owner_deltas.ipv6_routes, mtu: mtu_delta(old_mtu, self.effective_mtu()), - }) + } } pub fn effective_mtu(&self) -> Option { @@ -784,32 +797,13 @@ 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 + SharedIfConfigClaims { + ipv4_addresses: merged_items(&old_claims.ipv4_addresses, &next_claims.ipv4_addresses), + ipv6_addresses: merged_items(&old_claims.ipv6_addresses, &next_claims.ipv6_addresses), + ipv4_routes: merged_items(&old_claims.ipv4_routes, &next_claims.ipv4_routes), + ipv6_routes: merged_items(&old_claims.ipv6_routes, &next_claims.ipv6_routes), + mtu: None, + } } async fn add_shared_ipv4_route( @@ -849,6 +843,27 @@ fn old_ipv4_route_hints( .collect() } +fn merged_items(old_items: &BTreeSet, new_items: &BTreeSet) -> BTreeSet +where + T: Ord + Clone, +{ + old_items.union(new_items).cloned().collect() +} + +fn remove_claimed_item(items: &mut BTreeSet, item: Option) +where + T: Ord, +{ + match item { + Some(item) => { + items.remove(&item); + } + None => { + items.clear(); + } + } +} + fn ignore_removed_ifcfg_not_found(result: Result<(), Error>) -> Result<(), Error> { match result { Err(Error::NotFound) => Ok(()), @@ -1031,13 +1046,8 @@ impl SharedVirtualNicMember { } pub async fn remove_ip(&self, ip: Option) -> Result<(), Error> { - self.update_claims(|claims| match ip { - Some(ip) => { - claims.ipv4_addresses.remove(&ip); - } - None => { - claims.ipv4_addresses.clear(); - } + self.update_claims(|claims| { + remove_claimed_item(&mut claims.ipv4_addresses, ip); }) .await } @@ -1083,13 +1093,8 @@ impl SharedVirtualNicMember { } pub async fn remove_ipv6(&self, ip: Option) -> Result<(), Error> { - self.update_claims(|claims| match ip { - Some(ip) => { - claims.ipv6_addresses.remove(&ip); - } - None => { - claims.ipv6_addresses.clear(); - } + self.update_claims(|claims| { + remove_claimed_item(&mut claims.ipv6_addresses, ip); }) .await } diff --git a/easytier/src/instance/shared_virtual_nic/dispatcher.rs b/easytier/src/instance/shared_virtual_nic/dispatcher.rs index 51032711..5010d167 100644 --- a/easytier/src/instance/shared_virtual_nic/dispatcher.rs +++ b/easytier/src/instance/shared_virtual_nic/dispatcher.rs @@ -168,11 +168,12 @@ impl SharedVirtualNicMemberTunnelTable { } } - let _ = reader_control_sender.send(SharedVirtualNicControl::Unregister { + notify_member_tunnel_closed( + &reader_control_sender, + &reader_close_notifier, member_id, registration_id, - }); - reader_close_notifier.notify_one(); + ); })); let writer_control_sender = control_sender.clone(); @@ -181,11 +182,12 @@ 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 { + notify_member_tunnel_closed( + &writer_control_sender, + &writer_close_notifier, member_id, registration_id, - }); - writer_close_notifier.notify_one(); + ); break; } } @@ -234,6 +236,19 @@ impl SharedVirtualNicMemberTunnelTable { } } +fn notify_member_tunnel_closed( + control_sender: &mpsc::UnboundedSender, + close_notifier: &Notify, + member_id: SharedVirtualNicMemberId, + registration_id: SharedVirtualNicMemberRegistrationId, +) { + let _ = control_sender.send(SharedVirtualNicControl::Unregister { + member_id, + registration_id, + }); + close_notifier.notify_one(); +} + #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] enum SharedVirtualNicFlowAddr { V4(u32), @@ -609,27 +624,31 @@ impl SharedVirtualNicDispatcher { member_id: SharedVirtualNicMemberId, claims: &SharedIfConfigClaims, ) -> Result<(), Error> { - let (ack, rx) = oneshot::channel(); - self.control_sender - .send(SharedVirtualNicControl::UpdateSources { - member_id, - sources: SharedVirtualNicMemberSources::from_claims(claims), - ack, - }) - .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; - rx.await - .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running").into()) + self.send_source_update( + member_id, + SharedVirtualNicMemberSources::from_claims(claims), + ) + .await } pub(super) async fn remove_sources( &self, member_id: SharedVirtualNicMemberId, + ) -> Result<(), Error> { + self.send_source_update(member_id, SharedVirtualNicMemberSources::default()) + .await + } + + async fn send_source_update( + &self, + member_id: SharedVirtualNicMemberId, + sources: SharedVirtualNicMemberSources, ) -> Result<(), Error> { let (ack, rx) = oneshot::channel(); self.control_sender .send(SharedVirtualNicControl::UpdateSources { member_id, - sources: SharedVirtualNicMemberSources::default(), + sources, ack, }) .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; @@ -659,6 +678,43 @@ impl SharedVirtualNicDispatcher { } } +enum DispatcherControlResult { + Continue, + Stop { + invalidate: bool, + ack: Option>, + }, +} + +fn handle_dispatcher_control( + state: &mut SharedVirtualNicDispatcherState, + control: Option, +) -> DispatcherControlResult { + let Some(control) = control else { + return DispatcherControlResult::Stop { + invalidate: true, + ack: None, + }; + }; + + match control { + SharedVirtualNicControl::Shutdown { invalidate, ack } => DispatcherControlResult::Stop { + invalidate, + ack: Some(ack), + }, + other => { + state.handle_control(other); + DispatcherControlResult::Continue + } + } +} + +fn acknowledge_dispatcher_shutdown(ack: Option>) { + if let Some(ack) = ack { + let _ = ack.send(()); + } +} + struct SharedVirtualNicDispatcherTask { tun_stream: Pin>, tun_sink: Pin>, @@ -674,16 +730,12 @@ impl SharedVirtualNicDispatcherTask { loop { tokio::select! { control = self.control_receiver.recv() => { - let Some(control) = control else { - break; - }; - match control { - SharedVirtualNicControl::Shutdown { invalidate, ack } => { - self.cleanup(invalidate); - let _ = ack.send(()); - return; - } - other => self.state.handle_control(other), + if let DispatcherControlResult::Stop { invalidate, ack } = + handle_dispatcher_control(&mut self.state, control) + { + self.cleanup(invalidate); + acknowledge_dispatcher_shutdown(ack); + return; } } member_packet = self.to_tun_receiver.recv() => { @@ -900,21 +952,13 @@ impl SharedVirtualNicMobileDispatcherTask { } fn handle_control(&mut self, control: Option) -> bool { - let Some(control) = control else { - self.cleanup(true); - return false; - }; - - match control { - SharedVirtualNicControl::Shutdown { invalidate, ack } => { + match handle_dispatcher_control(&mut self.state, control) { + DispatcherControlResult::Continue => true, + DispatcherControlResult::Stop { invalidate, ack } => { self.cleanup(invalidate); - let _ = ack.send(()); + acknowledge_dispatcher_shutdown(ack); false } - other => { - self.state.handle_control(other); - true - } } } diff --git a/easytier/src/instance/virtual_nic.rs b/easytier/src/instance/virtual_nic.rs index 3b38e775..dd8f3a5d 100644 --- a/easytier/src/instance/virtual_nic.rs +++ b/easytier/src/instance/virtual_nic.rs @@ -1318,6 +1318,13 @@ impl NicCtx { self.backend.ifname().await } + async fn tun_ifname(&self) -> Result { + self.backend + .ifname() + .await + .ok_or_else(|| anyhow::anyhow!("tun device has no interface name").into()) + } + pub async fn assign_ipv4_to_tun_device(&self, ipv4_addr: cidr::Ipv4Inet) -> Result<(), Error> { self.backend.link_up().await?; self.backend.remove_ip(None).await?; @@ -1498,6 +1505,15 @@ impl NicCtx { }); } + fn start_tunnel_forwarding(&mut self, tunnel: Box) -> Result<(), Error> { + let (stream, sink) = tunnel.split(); + + self.do_forward_nic_to_peers_task(stream)?; + self.do_forward_peers_to_nic(sink); + + Ok(()) + } + #[cfg(target_os = "windows")] fn start_windows_udp_broadcast_relay(&mut self, virtual_ipv4: Ipv4Inet) { if !self.global_ctx.get_flags().enable_udp_broadcast_relay { @@ -1792,11 +1808,7 @@ impl NicCtx { ) -> Result<(), Error> { let tunnel = match self.backend.create_dev().await { Ok(ret) => { - let ifname = self - .backend - .ifname() - .await - .ok_or_else(|| anyhow::anyhow!("tun device has no interface name"))?; + let ifname = self.tun_ifname().await?; #[cfg(target_os = "windows")] { @@ -1831,10 +1843,7 @@ impl NicCtx { } }; - let (stream, sink) = tunnel.split(); - - self.do_forward_nic_to_peers_task(stream)?; - self.do_forward_peers_to_nic(sink); + self.start_tunnel_forwarding(tunnel)?; // Assign IPv4 address if provided if let Some(ipv4_addr) = ipv4_addr { @@ -1865,11 +1874,7 @@ impl NicCtx { ) -> Result<(), Error> { 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"))?; + let ifname = self.tun_ifname().await?; (ret, ifname) } Err(err) => { @@ -1899,10 +1904,7 @@ impl NicCtx { self.global_ctx .issue_event(GlobalCtxEvent::TunDeviceReady(ifname)); - let (stream, sink) = tunnel.split(); - - self.do_forward_nic_to_peers_task(stream)?; - self.do_forward_peers_to_nic(sink); + self.start_tunnel_forwarding(tunnel)?; Ok(()) }