mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-08-06 20:49:46 +00:00
refactor: keep shared dispatcher state local
Move shared member and flow ownership state into the dispatcher task. Use control messages for member register and unregister events. Keep the member table lock off the packet forwarding path.
This commit is contained in:
@@ -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<StdMutex<SharedVirtualNicMemberTunnelTableState>>,
|
||||
@@ -37,8 +47,8 @@ pub(super) struct SharedVirtualNicMemberTunnelTable {
|
||||
|
||||
#[derive(Default)]
|
||||
struct SharedVirtualNicMemberTunnelTableState {
|
||||
members: BTreeMap<SharedVirtualNicMemberId, SharedVirtualNicMemberTunnelEntry>,
|
||||
to_tun_sender: Option<mpsc::Sender<SharedVirtualNicMemberPacket>>,
|
||||
control_sender: Option<mpsc::UnboundedSender<SharedVirtualNicControl>>,
|
||||
}
|
||||
|
||||
struct SharedVirtualNicMemberTunnelEntry {
|
||||
@@ -48,8 +58,20 @@ struct SharedVirtualNicMemberTunnelEntry {
|
||||
}
|
||||
|
||||
impl SharedVirtualNicMemberTunnelTable {
|
||||
fn attach_dispatcher(&self, sender: mpsc::Sender<SharedVirtualNicMemberPacket>) {
|
||||
self.state.lock().unwrap().to_tun_sender = Some(sender);
|
||||
fn attach_dispatcher(
|
||||
&self,
|
||||
to_tun_sender: mpsc::Sender<SharedVirtualNicMemberPacket>,
|
||||
control_sender: mpsc::UnboundedSender<SharedVirtualNicControl>,
|
||||
) {
|
||||
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<dyn Tunnel>,
|
||||
close_notifier: Arc<Notify>,
|
||||
) -> 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<SharedVirtualNicMemberId>,
|
||||
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<SharedVirtualNicMemberPacket>,
|
||||
mpsc::UnboundedSender<SharedVirtualNicControl>,
|
||||
)> {
|
||||
let state = self.state.lock().unwrap();
|
||||
Some((state.to_tun_sender.clone()?, state.control_sender.clone()?))
|
||||
}
|
||||
|
||||
fn member_sender(&self, member_id: SharedVirtualNicMemberId) -> Option<mpsc::Sender<ZCPacket>> {
|
||||
self.state
|
||||
.lock()
|
||||
.unwrap()
|
||||
.members
|
||||
.get(&member_id)
|
||||
.map(|entry| entry.sender.clone())
|
||||
}
|
||||
|
||||
fn first_member_sender(&self) -> Option<(SharedVirtualNicMemberId, mpsc::Sender<ZCPacket>)> {
|
||||
self.state
|
||||
.lock()
|
||||
.unwrap()
|
||||
.members
|
||||
.iter()
|
||||
.next()
|
||||
.map(|(member_id, entry)| (*member_id, entry.sender.clone()))
|
||||
fn control_sender(&self) -> Option<mpsc::UnboundedSender<SharedVirtualNicControl>> {
|
||||
self.state.lock().unwrap().control_sender.clone()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -246,35 +221,38 @@ impl SharedVirtualNicFlowKey {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
#[derive(Default)]
|
||||
struct SharedVirtualNicFlowTable {
|
||||
owners: Arc<StdMutex<BTreeMap<SharedVirtualNicFlowKey, SharedVirtualNicMemberId>>>,
|
||||
owners: BTreeMap<SharedVirtualNicFlowKey, SharedVirtualNicMemberId>,
|
||||
}
|
||||
|
||||
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<SharedVirtualNicMemberId> {
|
||||
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<AbortOnDropHandle<()>>,
|
||||
_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<SharedVirtualNicMemberPacket>,
|
||||
mut tun_sink: Pin<Box<dyn ZCPacketSink>>,
|
||||
flow_table: SharedVirtualNicFlowTable,
|
||||
member_tunnel_table: SharedVirtualNicMemberTunnelTable,
|
||||
valid: Arc<AtomicBool>,
|
||||
) {
|
||||
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<Box<dyn ZCPacketStream>>,
|
||||
member_tunnel_table: SharedVirtualNicMemberTunnelTable,
|
||||
flow_table: SharedVirtualNicFlowTable,
|
||||
valid: Arc<AtomicBool>,
|
||||
) {
|
||||
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<Box<dyn ZCPacketStream>>,
|
||||
tun_sink: Pin<Box<dyn ZCPacketSink>>,
|
||||
to_tun_receiver: mpsc::Receiver<SharedVirtualNicMemberPacket>,
|
||||
control_receiver: mpsc::UnboundedReceiver<SharedVirtualNicControl>,
|
||||
member_tunnel_table: SharedVirtualNicMemberTunnelTable,
|
||||
valid: Arc<AtomicBool>,
|
||||
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<SharedVirtualNicMemberId, SharedVirtualNicMemberTunnelEntry>,
|
||||
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<SharedVirtualNicMemberId>,
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user