refactor: simplify shared virtual nic helpers

This commit is contained in:
sijie.sun
2026-06-16 12:46:28 +08:00
parent 800c840bb5
commit d742fa34aa
4 changed files with 269 additions and 203 deletions
+42 -27
View File
@@ -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<PeerManager>,
nic_ctx: NicCtx,
tun_dev: Option<String>,
tun_ip: Option<Ipv4Inet>,
) {
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;
+124 -119
View File
@@ -103,6 +103,13 @@ pub struct SharedIfConfigDelta {
pub mtu: Option<SharedMtuChange>,
}
struct SharedIfConfigOwnerDeltas {
ipv4_addresses: OwnedItemDelta<Ipv4Inet>,
ipv6_addresses: OwnedItemDelta<Ipv6Inet>,
ipv4_routes: OwnedItemDelta<SharedIpv4Route>,
ipv6_routes: OwnedItemDelta<SharedIpv6Route>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct SharedIfConfigSnapshot {
pub ipv4_addresses: BTreeMap<Ipv4Inet, BTreeSet<SharedVirtualNicMemberId>>,
@@ -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::<BTreeSet<_>>();
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<u32>,
old_ipv4_route_sources: &BTreeMap<SharedIpv4Route, Option<Ipv4Addr>>,
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<u32> {
@@ -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<T>(old_items: &BTreeSet<T>, new_items: &BTreeSet<T>) -> BTreeSet<T>
where
T: Ord + Clone,
{
old_items.union(new_items).cloned().collect()
}
fn remove_claimed_item<T>(items: &mut BTreeSet<T>, item: Option<T>)
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<Ipv4Inet>) -> 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<Ipv6Inet>) -> 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
}
@@ -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<SharedVirtualNicControl>,
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<oneshot::Sender<()>>,
},
}
fn handle_dispatcher_control(
state: &mut SharedVirtualNicDispatcherState,
control: Option<SharedVirtualNicControl>,
) -> 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<oneshot::Sender<()>>) {
if let Some(ack) = ack {
let _ = ack.send(());
}
}
struct SharedVirtualNicDispatcherTask {
tun_stream: Pin<Box<dyn ZCPacketStream>>,
tun_sink: Pin<Box<dyn ZCPacketSink>>,
@@ -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<SharedVirtualNicControl>) -> 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
}
}
}
+20 -18
View File
@@ -1318,6 +1318,13 @@ impl NicCtx {
self.backend.ifname().await
}
async fn tun_ifname(&self) -> Result<String, Error> {
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<dyn Tunnel>) -> 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(())
}