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:
sijie.sun
2026-06-14 01:26:19 +08:00
parent b2f1b37336
commit be0859aca6
@@ -12,7 +12,7 @@ use futures::{SinkExt, StreamExt};
use pnet::packet::{ use pnet::packet::{
Packet as _, ipv4::Ipv4Packet, ipv6::Ipv6Packet, tcp::TcpPacket, udp::UdpPacket, 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 tokio_util::task::AbortOnDropHandle;
use crate::{ use crate::{
@@ -30,6 +30,16 @@ struct SharedVirtualNicMemberPacket {
packet: ZCPacket, packet: ZCPacket,
} }
enum SharedVirtualNicControl {
Register {
member_id: SharedVirtualNicMemberId,
entry: SharedVirtualNicMemberTunnelEntry,
},
Unregister {
member_id: SharedVirtualNicMemberId,
},
}
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub(super) struct SharedVirtualNicMemberTunnelTable { pub(super) struct SharedVirtualNicMemberTunnelTable {
state: Arc<StdMutex<SharedVirtualNicMemberTunnelTableState>>, state: Arc<StdMutex<SharedVirtualNicMemberTunnelTableState>>,
@@ -37,8 +47,8 @@ pub(super) struct SharedVirtualNicMemberTunnelTable {
#[derive(Default)] #[derive(Default)]
struct SharedVirtualNicMemberTunnelTableState { struct SharedVirtualNicMemberTunnelTableState {
members: BTreeMap<SharedVirtualNicMemberId, SharedVirtualNicMemberTunnelEntry>,
to_tun_sender: Option<mpsc::Sender<SharedVirtualNicMemberPacket>>, to_tun_sender: Option<mpsc::Sender<SharedVirtualNicMemberPacket>>,
control_sender: Option<mpsc::UnboundedSender<SharedVirtualNicControl>>,
} }
struct SharedVirtualNicMemberTunnelEntry { struct SharedVirtualNicMemberTunnelEntry {
@@ -48,8 +58,20 @@ struct SharedVirtualNicMemberTunnelEntry {
} }
impl SharedVirtualNicMemberTunnelTable { impl SharedVirtualNicMemberTunnelTable {
fn attach_dispatcher(&self, sender: mpsc::Sender<SharedVirtualNicMemberPacket>) { fn attach_dispatcher(
self.state.lock().unwrap().to_tun_sender = Some(sender); &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( pub(super) fn register(
@@ -58,20 +80,21 @@ impl SharedVirtualNicMemberTunnelTable {
tunnel: Box<dyn Tunnel>, tunnel: Box<dyn Tunnel>,
close_notifier: Arc<Notify>, close_notifier: Arc<Notify>,
) -> Result<(), Error> { ) -> Result<(), Error> {
let to_tun_sender = self let channels = self
.state .dispatcher_channels()
.lock()
.unwrap()
.to_tun_sender
.clone()
.ok_or_else(|| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?; .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 (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_close_notifier = close_notifier.clone();
let reader_task = AbortOnDropHandle::new(tokio::spawn(async move { 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 { while let Some(packet) = member_stream.next().await {
let packet = match packet { let packet = match packet {
Ok(packet) => 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(); 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_close_notifier = close_notifier.clone();
let writer_task = AbortOnDropHandle::new(tokio::spawn(async move { 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 { if let Err(err) = member_sink.send(packet).await {
tracing::error!(?member_id, ?err, "shared member tunnel write failed"); 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(); writer_close_notifier.notify_one();
break; break;
} }
@@ -108,86 +132,37 @@ impl SharedVirtualNicMemberTunnelTable {
})); }));
let entry = SharedVirtualNicMemberTunnelEntry { let entry = SharedVirtualNicMemberTunnelEntry {
sender, sender: to_member_sender,
close_notifier, close_notifier,
_tasks: vec![reader_task, writer_task], _tasks: vec![reader_task, writer_task],
}; };
let old_entry = { control_sender
let mut state = self.state.lock().unwrap(); .send(SharedVirtualNicControl::Register { member_id, entry })
state.members.insert(member_id, entry) .map_err(|_| anyhow::anyhow!("shared virtual nic dispatcher is not running"))?;
}; let _ = reader_start_sender.send(());
drop(old_entry);
Ok(()) Ok(())
} }
pub(super) fn unregister(&self, member_id: SharedVirtualNicMemberId) { pub(super) fn unregister(&self, member_id: SharedVirtualNicMemberId) {
let entry = { let Some(control_sender) = self.control_sender() else {
let mut state = self.state.lock().unwrap(); return;
state.members.remove(&member_id)
}; };
drop(entry); let _ = control_sender.send(SharedVirtualNicControl::Unregister { member_id });
} }
fn close_all(&self) { fn dispatcher_channels(
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(
&self, &self,
preferred_member_id: Option<SharedVirtualNicMemberId>, ) -> Option<(
packet: ZCPacket, mpsc::Sender<SharedVirtualNicMemberPacket>,
) -> bool { mpsc::UnboundedSender<SharedVirtualNicControl>,
let mut packet = packet; )> {
if let Some(member_id) = preferred_member_id { let state = self.state.lock().unwrap();
if let Some(sender) = self.member_sender(member_id) { Some((state.to_tun_sender.clone()?, state.control_sender.clone()?))
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
}
}
} }
fn member_sender(&self, member_id: SharedVirtualNicMemberId) -> Option<mpsc::Sender<ZCPacket>> { fn control_sender(&self) -> Option<mpsc::UnboundedSender<SharedVirtualNicControl>> {
self.state self.state.lock().unwrap().control_sender.clone()
.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()))
} }
} }
@@ -246,35 +221,38 @@ impl SharedVirtualNicFlowKey {
} }
} }
#[derive(Clone, Default)] #[derive(Default)]
struct SharedVirtualNicFlowTable { struct SharedVirtualNicFlowTable {
owners: Arc<StdMutex<BTreeMap<SharedVirtualNicFlowKey, SharedVirtualNicMemberId>>>, owners: BTreeMap<SharedVirtualNicFlowKey, SharedVirtualNicMemberId>,
} }
impl SharedVirtualNicFlowTable { 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()) let Some(key) = SharedVirtualNicFlowKey::from_packet(packet).map(|key| key.reversed())
else { else {
return; return;
}; };
let mut owners = self.owners.lock().unwrap(); if !self.owners.contains_key(&key) && self.owners.len() >= FLOW_OWNER_LIMIT {
if !owners.contains_key(&key) && owners.len() >= FLOW_OWNER_LIMIT { if let Some(oldest_key) = self.owners.keys().next().cloned() {
if let Some(oldest_key) = owners.keys().next().cloned() { self.owners.remove(&oldest_key);
owners.remove(&oldest_key);
} }
} }
owners.insert(key, member_id); self.owners.insert(key, member_id);
} }
fn owner_of(&self, packet: &ZCPacket) -> Option<SharedVirtualNicMemberId> { fn owner_of(&self, packet: &ZCPacket) -> Option<SharedVirtualNicMemberId> {
let key = SharedVirtualNicFlowKey::from_packet(packet)?; 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 { pub(super) struct SharedVirtualNicDispatcher {
_tasks: Vec<AbortOnDropHandle<()>>, _task: AbortOnDropHandle<()>,
} }
impl SharedVirtualNicDispatcher { impl SharedVirtualNicDispatcher {
@@ -285,69 +263,183 @@ impl SharedVirtualNicDispatcher {
) -> Self { ) -> Self {
let (tun_stream, tun_sink) = tunnel.split(); let (tun_stream, tun_sink) = tunnel.split();
let (to_tun_sender, to_tun_receiver) = mpsc::channel(MEMBER_TUNNEL_BUFFER_SIZE); 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 task = SharedVirtualNicDispatcherTask {
let tasks = vec![ tun_stream,
AbortOnDropHandle::new(tokio::spawn(Self::forward_members_to_tun( tun_sink,
to_tun_receiver, to_tun_receiver,
tun_sink, control_receiver,
flow_table.clone(), member_tunnel_table,
member_tunnel_table.clone(), valid,
valid.clone(), state: SharedVirtualNicDispatcherState::default(),
))), };
AbortOnDropHandle::new(tokio::spawn(Self::forward_tun_to_members(
tun_stream,
member_tunnel_table,
flow_table,
valid,
))),
];
Self { _tasks: tasks } Self {
} _task: AbortOnDropHandle::new(tokio::spawn(task.run())),
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;
}
} }
valid.store(false, Ordering::Release);
member_tunnel_table.close_all();
} }
}
async fn forward_tun_to_members( struct SharedVirtualNicDispatcherTask {
mut tun_stream: Pin<Box<dyn ZCPacketStream>>, tun_stream: Pin<Box<dyn ZCPacketStream>>,
member_tunnel_table: SharedVirtualNicMemberTunnelTable, tun_sink: Pin<Box<dyn ZCPacketSink>>,
flow_table: SharedVirtualNicFlowTable, to_tun_receiver: mpsc::Receiver<SharedVirtualNicMemberPacket>,
valid: Arc<AtomicBool>, control_receiver: mpsc::UnboundedReceiver<SharedVirtualNicControl>,
) { member_tunnel_table: SharedVirtualNicMemberTunnelTable,
while let Some(packet) = tun_stream.next().await { valid: Arc<AtomicBool>,
let packet = match packet { state: SharedVirtualNicDispatcherState,
Ok(packet) => packet, }
Err(err) => {
tracing::error!(?err, "shared virtual nic read from tun failed"); impl SharedVirtualNicDispatcherTask {
break; 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); self.valid.store(false, Ordering::Release);
member_tunnel_table.close_all(); 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)
}
}
} }
} }