diff --git a/easytier/src/instance/shared_virtual_nic/dispatcher.rs b/easytier/src/instance/shared_virtual_nic/dispatcher.rs index e2625c87..df5bd307 100644 --- a/easytier/src/instance/shared_virtual_nic/dispatcher.rs +++ b/easytier/src/instance/shared_virtual_nic/dispatcher.rs @@ -12,7 +12,7 @@ use futures::{SinkExt, StreamExt}; use pnet::packet::{ Packet as _, ipv4::Ipv4Packet, ipv6::Ipv6Packet, tcp::TcpPacket, udp::UdpPacket, }; -use tokio::sync::{Notify, mpsc}; +use tokio::sync::{Notify, mpsc, oneshot}; use tokio_util::task::AbortOnDropHandle; use crate::{ @@ -30,6 +30,16 @@ struct SharedVirtualNicMemberPacket { packet: ZCPacket, } +enum SharedVirtualNicControl { + Register { + member_id: SharedVirtualNicMemberId, + entry: SharedVirtualNicMemberTunnelEntry, + }, + Unregister { + member_id: SharedVirtualNicMemberId, + }, +} + #[derive(Clone, Default)] pub(super) struct SharedVirtualNicMemberTunnelTable { state: Arc>, @@ -37,8 +47,8 @@ pub(super) struct SharedVirtualNicMemberTunnelTable { #[derive(Default)] struct SharedVirtualNicMemberTunnelTableState { - members: BTreeMap, to_tun_sender: Option>, + control_sender: Option>, } struct SharedVirtualNicMemberTunnelEntry { @@ -48,8 +58,20 @@ struct SharedVirtualNicMemberTunnelEntry { } impl SharedVirtualNicMemberTunnelTable { - fn attach_dispatcher(&self, sender: mpsc::Sender) { - self.state.lock().unwrap().to_tun_sender = Some(sender); + fn attach_dispatcher( + &self, + to_tun_sender: mpsc::Sender, + control_sender: mpsc::UnboundedSender, + ) { + let mut state = self.state.lock().unwrap(); + state.to_tun_sender = Some(to_tun_sender); + state.control_sender = Some(control_sender); + } + + fn detach_dispatcher(&self) { + let mut state = self.state.lock().unwrap(); + state.to_tun_sender.take(); + state.control_sender.take(); } pub(super) fn register( @@ -58,20 +80,21 @@ impl SharedVirtualNicMemberTunnelTable { tunnel: Box, close_notifier: Arc, ) -> Result<(), Error> { - let to_tun_sender = self - .state - .lock() - .unwrap() - .to_tun_sender - .clone() + let channels = self + .dispatcher_channels() .ok_or_else(|| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; - + let (to_tun_sender, control_sender) = channels; let (mut member_stream, mut member_sink) = tunnel.split(); - let (sender, mut receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (to_member_sender, mut to_member_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); + let (reader_start_sender, reader_start_receiver) = oneshot::channel(); - let table = self.clone(); + let reader_control_sender = control_sender.clone(); let reader_close_notifier = close_notifier.clone(); let reader_task = AbortOnDropHandle::new(tokio::spawn(async move { + if reader_start_receiver.await.is_err() { + return; + } + while let Some(packet) = member_stream.next().await { let packet = match packet { Ok(packet) => packet, @@ -90,17 +113,18 @@ impl SharedVirtualNicMemberTunnelTable { } } - table.unregister(member_id); + let _ = reader_control_sender.send(SharedVirtualNicControl::Unregister { member_id }); reader_close_notifier.notify_one(); })); - let table = self.clone(); + let writer_control_sender = control_sender.clone(); let writer_close_notifier = close_notifier.clone(); let writer_task = AbortOnDropHandle::new(tokio::spawn(async move { - while let Some(packet) = receiver.recv().await { + 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"); - table.unregister(member_id); + let _ = writer_control_sender + .send(SharedVirtualNicControl::Unregister { member_id }); writer_close_notifier.notify_one(); break; } @@ -108,86 +132,37 @@ impl SharedVirtualNicMemberTunnelTable { })); let entry = SharedVirtualNicMemberTunnelEntry { - sender, + sender: to_member_sender, close_notifier, _tasks: vec![reader_task, writer_task], }; - let old_entry = { - let mut state = self.state.lock().unwrap(); - state.members.insert(member_id, entry) - }; - drop(old_entry); + control_sender + .send(SharedVirtualNicControl::Register { member_id, entry }) + .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; + let _ = reader_start_sender.send(()); + Ok(()) } pub(super) fn unregister(&self, member_id: SharedVirtualNicMemberId) { - let entry = { - let mut state = self.state.lock().unwrap(); - state.members.remove(&member_id) + let Some(control_sender) = self.control_sender() else { + return; }; - drop(entry); + let _ = control_sender.send(SharedVirtualNicControl::Unregister { member_id }); } - fn close_all(&self) { - let entries = { - let mut state = self.state.lock().unwrap(); - state.to_tun_sender.take(); - std::mem::take(&mut state.members) - }; - - for entry in entries.into_values() { - entry.close_notifier.notify_one(); - } - } - - async fn send_packet( + fn dispatcher_channels( &self, - preferred_member_id: Option, - packet: ZCPacket, - ) -> bool { - let mut packet = packet; - if let Some(member_id) = preferred_member_id { - if let Some(sender) = self.member_sender(member_id) { - match sender.send(packet).await { - Ok(()) => return true, - Err(err) => { - packet = err.0; - self.unregister(member_id); - } - } - } - } - - let Some((member_id, sender)) = self.first_member_sender() else { - return false; - }; - - match sender.send(packet).await { - Ok(()) => true, - Err(_) => { - self.unregister(member_id); - false - } - } + ) -> Option<( + mpsc::Sender, + mpsc::UnboundedSender, + )> { + let state = self.state.lock().unwrap(); + Some((state.to_tun_sender.clone()?, state.control_sender.clone()?)) } - fn member_sender(&self, member_id: SharedVirtualNicMemberId) -> Option> { - self.state - .lock() - .unwrap() - .members - .get(&member_id) - .map(|entry| entry.sender.clone()) - } - - fn first_member_sender(&self) -> Option<(SharedVirtualNicMemberId, mpsc::Sender)> { - self.state - .lock() - .unwrap() - .members - .iter() - .next() - .map(|(member_id, entry)| (*member_id, entry.sender.clone())) + fn control_sender(&self) -> Option> { + self.state.lock().unwrap().control_sender.clone() } } @@ -246,35 +221,38 @@ impl SharedVirtualNicFlowKey { } } -#[derive(Clone, Default)] +#[derive(Default)] struct SharedVirtualNicFlowTable { - owners: Arc>>, + owners: BTreeMap, } impl SharedVirtualNicFlowTable { - fn remember_reverse_owner(&self, member_id: SharedVirtualNicMemberId, packet: &ZCPacket) { + fn remember_reverse_owner(&mut self, member_id: SharedVirtualNicMemberId, packet: &ZCPacket) { let Some(key) = SharedVirtualNicFlowKey::from_packet(packet).map(|key| key.reversed()) else { return; }; - let mut owners = self.owners.lock().unwrap(); - if !owners.contains_key(&key) && owners.len() >= FLOW_OWNER_LIMIT { - if let Some(oldest_key) = owners.keys().next().cloned() { - owners.remove(&oldest_key); + if !self.owners.contains_key(&key) && self.owners.len() >= FLOW_OWNER_LIMIT { + if let Some(oldest_key) = self.owners.keys().next().cloned() { + self.owners.remove(&oldest_key); } } - owners.insert(key, member_id); + self.owners.insert(key, member_id); } fn owner_of(&self, packet: &ZCPacket) -> Option { let key = SharedVirtualNicFlowKey::from_packet(packet)?; - self.owners.lock().unwrap().get(&key).copied() + self.owners.get(&key).copied() + } + + fn remove_owner(&mut self, member_id: SharedVirtualNicMemberId) { + self.owners.retain(|_, owner| *owner != member_id); } } pub(super) struct SharedVirtualNicDispatcher { - _tasks: Vec>, + _task: AbortOnDropHandle<()>, } impl SharedVirtualNicDispatcher { @@ -285,69 +263,183 @@ impl SharedVirtualNicDispatcher { ) -> Self { let (tun_stream, tun_sink) = tunnel.split(); let (to_tun_sender, to_tun_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); - member_tunnel_table.attach_dispatcher(to_tun_sender); + let (control_sender, control_receiver) = mpsc::unbounded_channel(); + member_tunnel_table.attach_dispatcher(to_tun_sender, control_sender); - let flow_table = SharedVirtualNicFlowTable::default(); - let tasks = vec![ - AbortOnDropHandle::new(tokio::spawn(Self::forward_members_to_tun( - to_tun_receiver, - tun_sink, - flow_table.clone(), - member_tunnel_table.clone(), - valid.clone(), - ))), - AbortOnDropHandle::new(tokio::spawn(Self::forward_tun_to_members( - tun_stream, - member_tunnel_table, - flow_table, - valid, - ))), - ]; + let task = SharedVirtualNicDispatcherTask { + tun_stream, + tun_sink, + to_tun_receiver, + control_receiver, + member_tunnel_table, + valid, + state: SharedVirtualNicDispatcherState::default(), + }; - Self { _tasks: tasks } - } - - async fn forward_members_to_tun( - mut receiver: mpsc::Receiver, - mut tun_sink: Pin>, - flow_table: SharedVirtualNicFlowTable, - member_tunnel_table: SharedVirtualNicMemberTunnelTable, - valid: Arc, - ) { - while let Some(member_packet) = receiver.recv().await { - flow_table.remember_reverse_owner(member_packet.member_id, &member_packet.packet); - if let Err(err) = tun_sink.send(member_packet.packet).await { - tracing::error!(?err, "shared virtual nic write to tun failed"); - break; - } + Self { + _task: AbortOnDropHandle::new(tokio::spawn(task.run())), } - - valid.store(false, Ordering::Release); - member_tunnel_table.close_all(); } +} - async fn forward_tun_to_members( - mut tun_stream: Pin>, - member_tunnel_table: SharedVirtualNicMemberTunnelTable, - flow_table: SharedVirtualNicFlowTable, - valid: Arc, - ) { - while let Some(packet) = tun_stream.next().await { - let packet = match packet { - Ok(packet) => packet, - Err(err) => { - tracing::error!(?err, "shared virtual nic read from tun failed"); - break; +struct SharedVirtualNicDispatcherTask { + tun_stream: Pin>, + tun_sink: Pin>, + to_tun_receiver: mpsc::Receiver, + control_receiver: mpsc::UnboundedReceiver, + member_tunnel_table: SharedVirtualNicMemberTunnelTable, + valid: Arc, + state: SharedVirtualNicDispatcherState, +} + +impl SharedVirtualNicDispatcherTask { + async fn run(mut self) { + loop { + tokio::select! { + control = self.control_receiver.recv() => { + let Some(control) = control else { + break; + }; + self.state.handle_control(control); + } + member_packet = self.to_tun_receiver.recv() => { + let Some(member_packet) = member_packet else { + break; + }; + if !self.forward_member_packet_to_tun(member_packet).await { + break; + } + } + packet = self.tun_stream.next() => { + let Some(packet) = packet else { + break; + }; + let packet = match packet { + Ok(packet) => packet, + Err(err) => { + tracing::error!(?err, "shared virtual nic read from tun failed"); + break; + } + }; + self.state.forward_tun_packet_to_member(packet).await; } - }; - let member_id = flow_table.owner_of(&packet); - if !member_tunnel_table.send_packet(member_id, packet).await { - tracing::trace!("shared virtual nic dropped packet without active member"); } } - valid.store(false, Ordering::Release); - member_tunnel_table.close_all(); + self.valid.store(false, Ordering::Release); + self.member_tunnel_table.detach_dispatcher(); + self.state.close_all(); + } + + async fn forward_member_packet_to_tun( + &mut self, + member_packet: SharedVirtualNicMemberPacket, + ) -> bool { + self.state + .remember_reverse_owner(member_packet.member_id, &member_packet.packet); + if let Err(err) = self.tun_sink.send(member_packet.packet).await { + tracing::error!(?err, "shared virtual nic write to tun failed"); + return false; + } + true + } +} + +#[derive(Default)] +struct SharedVirtualNicDispatcherState { + members: BTreeMap, + flow_table: SharedVirtualNicFlowTable, +} + +impl SharedVirtualNicDispatcherState { + fn handle_control(&mut self, control: SharedVirtualNicControl) { + match control { + SharedVirtualNicControl::Register { member_id, entry } => { + self.register(member_id, entry); + } + SharedVirtualNicControl::Unregister { member_id } => { + self.unregister(member_id); + } + } + } + + fn register( + &mut self, + member_id: SharedVirtualNicMemberId, + entry: SharedVirtualNicMemberTunnelEntry, + ) { + let old_entry = self.members.insert(member_id, entry); + drop(old_entry); + } + + fn unregister(&mut self, member_id: SharedVirtualNicMemberId) { + let entry = self.members.remove(&member_id); + drop(entry); + self.flow_table.remove_owner(member_id); + } + + fn close_all(&mut self) { + let members = std::mem::take(&mut self.members); + self.flow_table.owners.clear(); + + for entry in members.into_values() { + entry.close_notifier.notify_one(); + } + } + + fn remember_reverse_owner(&mut self, member_id: SharedVirtualNicMemberId, packet: &ZCPacket) { + self.flow_table.remember_reverse_owner(member_id, packet); + } + + async fn forward_tun_packet_to_member(&mut self, packet: ZCPacket) { + let member_id = self.flow_table.owner_of(&packet); + if !self.send_packet(member_id, packet).await { + tracing::trace!("shared virtual nic dropped packet without active member"); + } + } + + async fn send_packet( + &mut self, + preferred_member_id: Option, + packet: ZCPacket, + ) -> bool { + let mut packet = packet; + if let Some(member_id) = preferred_member_id { + match self.send_packet_to_member(member_id, packet).await { + Ok(()) => return true, + Err(packet_on_failure) => { + packet = packet_on_failure; + } + } + } + + let Some(member_id) = self.members.keys().next().copied() else { + return false; + }; + + self.send_packet_to_member(member_id, packet).await.is_ok() + } + + async fn send_packet_to_member( + &mut self, + member_id: SharedVirtualNicMemberId, + packet: ZCPacket, + ) -> Result<(), ZCPacket> { + let Some(sender) = self + .members + .get(&member_id) + .map(|entry| entry.sender.clone()) + else { + return Err(packet); + }; + + match sender.send(packet).await { + Ok(()) => Ok(()), + Err(err) => { + self.unregister(member_id); + Err(err.0) + } + } } }