From b0d6b1f88429e040470c11790cfe098981aae1de Mon Sep 17 00:00:00 2001 From: lubeilin <1791778603@qq.com> Date: Fri, 1 Mar 2024 23:40:14 +0800 Subject: [PATCH] =?UTF-8?q?[mio]=20=E7=AE=80=E5=8C=96=E6=B5=81=E9=87=8F?= =?UTF-8?q?=E7=BB=9F=E8=AE=A1=EF=BC=8C=E5=85=BC=E5=AE=B932=E4=BD=8D?= =?UTF-8?q?=E7=B3=BB=E7=BB=9F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- vnt/src/core/conn.rs | 15 ++-- vnt/src/handle/tun_tap/tun_handler.rs | 23 ++---- vnt/src/util/counter/adder.rs | 109 +++++++++++++++++--------- 3 files changed, 88 insertions(+), 59 deletions(-) diff --git a/vnt/src/core/conn.rs b/vnt/src/core/conn.rs index 9e1ddcb..9b116b2 100644 --- a/vnt/src/core/conn.rs +++ b/vnt/src/core/conn.rs @@ -1,7 +1,6 @@ use std::collections::HashMap; use std::io; use std::net::Ipv4Addr; -use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::mpsc::{sync_channel, Receiver}; use std::sync::Arc; use std::time::Duration; @@ -26,7 +25,9 @@ use crate::handle::{ maintain, tun_tap, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, }; use crate::nat::NatTest; -use crate::util::{Scheduler, StopManager, U64Adder, WatchU64Adder}; +use crate::util::{ + Scheduler, SingleU64Adder, StopManager, U64Adder, WatchSingleU64Adder, WatchU64Adder, +}; use crate::{nat, tun_tap_device, DeviceInfo, VntCallback}; #[derive(Clone)] @@ -39,7 +40,7 @@ pub struct Vnt { context: Context, peer_nat_info_map: Arc>>, down_count_watcher: WatchU64Adder, - up_count_watcher: Arc, + up_count_watcher: WatchSingleU64Adder, } impl Vnt { @@ -135,7 +136,7 @@ impl Vnt { let (punch_sender, punch_receiver) = sync_channel(3); let peer_nat_info_map: Arc>> = Arc::new(RwLock::new(HashMap::with_capacity(16))); - let down_counter = U64Adder::with_capacity(8); + let down_counter = U64Adder::with_capacity(config.ports.as_ref().map(|v| v.len()).unwrap_or_default() + 8); let down_count_watcher = down_counter.watch(); let handler = RecvDataHandler::new( #[cfg(feature = "server_encrypt")] @@ -168,8 +169,8 @@ impl Vnt { config.tcp, tcp_socket_sender.clone(), ); - let up_counter = Arc::new(AtomicU64::new(0)); - let up_count_watcher = up_counter.clone(); + let up_counter = SingleU64Adder::new(); + let up_count_watcher = up_counter.watch(); tun_tap::tun_handler::start( stop_manager.clone(), context.clone(), @@ -345,7 +346,7 @@ impl Vnt { self.context.route_table.route_table() } pub fn up_stream(&self) -> u64 { - self.up_count_watcher.load(Ordering::Relaxed) + self.up_count_watcher.get() } pub fn down_stream(&self) -> u64 { self.down_count_watcher.get() diff --git a/vnt/src/handle/tun_tap/tun_handler.rs b/vnt/src/handle/tun_tap/tun_handler.rs index fe03108..b7ae2e3 100644 --- a/vnt/src/handle/tun_tap/tun_handler.rs +++ b/vnt/src/handle/tun_tap/tun_handler.rs @@ -1,4 +1,3 @@ -use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; use std::{io, thread}; @@ -18,7 +17,7 @@ use crate::handle::tun_tap::channel_group::{channel_group, GroupSyncSender}; use crate::handle::CurrentDeviceInfo; #[cfg(feature = "ip_proxy")] use crate::ip_proxy::IpProxyMap; -use crate::util::StopManager; +use crate::util::{SingleU64Adder, StopManager}; fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::Result<()> { if ipv4_packet.protocol() == ipv4::protocol::Protocol::Icmp { @@ -81,7 +80,7 @@ pub fn start( client_cipher: Cipher, server_cipher: Cipher, parallel: usize, - up_counter: Arc, + mut up_counter: SingleU64Adder, ) -> io::Result<()> { let worker = { let device = device.clone(); @@ -133,7 +132,7 @@ pub fn start( thread::Builder::new() .name("tun_handler".into()) .spawn(move || { - if let Err(e) = start_multi(stop_manager, device, sender, &up_counter) { + if let Err(e) = start_multi(stop_manager, device, sender, &mut up_counter) { log::warn!("stop:{}", e); } worker.stop_all(); @@ -152,7 +151,7 @@ pub fn start( ip_proxy_map, client_cipher, server_cipher, - &up_counter, + &mut up_counter, ) { log::warn!("stop:{}", e); } @@ -171,7 +170,7 @@ fn start_simple( #[cfg(feature = "ip_proxy")] ip_proxy_map: Option, client_cipher: Cipher, server_cipher: Cipher, - up_counter: &AtomicU64, + up_counter: &mut SingleU64Adder, ) -> io::Result<()> { let mut buf = [0; 1024 * 16]; loop { @@ -180,10 +179,7 @@ fn start_simple( } let len = device.read(&mut buf[12..])? + 12; //单线程的 - up_counter.store( - up_counter.load(Ordering::Relaxed) + len as u64, - Ordering::Relaxed, - ); + up_counter.add(len as u64); #[cfg(any(target_os = "macos"))] let mut buf = &mut buf[4..]; // buf是重复利用的,需要重置头部 @@ -212,7 +208,7 @@ fn start_multi( stop_manager: StopManager, device: Arc, mut group_sync_sender: GroupSyncSender<(Vec, usize)>, - up_counter: &AtomicU64, + up_counter: &mut SingleU64Adder, ) -> io::Result<()> { loop { if stop_manager.is_stop() { @@ -221,10 +217,7 @@ fn start_multi( let mut buf = vec![0; 1024 * 16]; let len = device.read(&mut buf[12..])? + 12; //单线程的 - up_counter.store( - up_counter.load(Ordering::Relaxed) + len as u64, - Ordering::Relaxed, - ); + up_counter.add(len as u64); if group_sync_sender.send((buf, len)).is_err() { return Ok(()); } diff --git a/vnt/src/util/counter/adder.rs b/vnt/src/util/counter/adder.rs index e539969..fc150e3 100644 --- a/vnt/src/util/counter/adder.rs +++ b/vnt/src/util/counter/adder.rs @@ -1,25 +1,68 @@ -use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; -/// 并发计数器,销毁计数器并不会释放计数槽,这不适用计数器会多次创建销毁的场景 +/// 不安全的并发计数器,谨慎使用 + pub struct U64Adder { + global_index: Arc, inner: Arc, - index: Option, + index: usize, +} +pub struct SingleU64Adder { + inner: Arc, +} +impl SingleU64Adder { + pub fn new() -> Self { + Self { + inner: Arc::new(SingleU64AdderInner::new()), + } + } + pub fn add(&mut self, num: u64) { + self.inner.add(num); + } + pub fn get(&self) -> u64 { + self.inner.get() + } + pub fn watch(&self) -> WatchSingleU64Adder { + WatchSingleU64Adder { + inner: self.inner.clone(), + } + } } +struct SingleU64AdderInner { + ptr: *mut u64, +} + +impl SingleU64AdderInner { + fn new() -> Self { + Self { + ptr: Box::into_raw(Box::new(0)), + } + } + #[inline(always)] + fn add(&self, num: u64) { + unsafe { *self.ptr += num } + } + + fn get(&self) -> u64 { + unsafe { *self.ptr } + } +} + +unsafe impl Send for SingleU64AdderInner {} + +unsafe impl Sync for SingleU64AdderInner {} + struct U64AdderInner { - global: AtomicU64, - base: Vec, + base: Vec, } impl U64AdderInner { pub fn get(&self) -> u64 { - let mut count = self.global.load(Ordering::Relaxed); + let mut count = 0; for counter in self.base.iter() { - let num = counter.load(Ordering::Relaxed); - if num > 1 { - count = count + num - 1; - } + count += counter.get() } count } @@ -29,27 +72,18 @@ impl U64Adder { /// 计数槽容量 pub fn with_capacity(capacity: usize) -> Self { let mut base = Vec::with_capacity(capacity); - base.push(AtomicU64::new(1)); - for _ in 1..capacity { - base.push(AtomicU64::new(0)) + for _ in 0..capacity { + base.push(SingleU64AdderInner::new()) } - let inner = Arc::new(U64AdderInner { - global: AtomicU64::new(0), - base, - }); + let inner = Arc::new(U64AdderInner { base }); U64Adder { + global_index: Arc::new(AtomicUsize::new(1)), inner, - index: Some(0), + index: 0, } } pub fn add(&mut self, num: u64) { - if let Some(index) = self.index { - let counter = &self.inner.base[index]; - let i = counter.load(Ordering::Relaxed); - counter.store(i + num, Ordering::Relaxed); - } else { - self.inner.global.fetch_add(num, Ordering::Relaxed); - } + self.inner.base[self.index].add(num); } pub fn get(&self) -> u64 { self.inner.get() @@ -63,21 +97,13 @@ impl U64Adder { impl Clone for U64Adder { fn clone(&self) -> Self { - let mut index: Option = None; - for (i, counter) in self.inner.base.iter().enumerate() { - //占用一个空闲的计数槽 - if counter.load(Ordering::Acquire) == 0 { - if counter - .compare_exchange(0, 1, Ordering::AcqRel, Ordering::Relaxed) - .is_ok() - { - index = Some(i); - break; - } - } + let index = self.global_index.fetch_add(1, Ordering::AcqRel); + if index > self.inner.base.len() { + panic!() } Self { + global_index: self.global_index.clone(), inner: self.inner.clone(), index, } @@ -94,3 +120,12 @@ impl WatchU64Adder { self.inner.get() } } +#[derive(Clone)] +pub struct WatchSingleU64Adder { + inner: Arc, +} +impl WatchSingleU64Adder { + pub fn get(&self) -> u64 { + self.inner.get() + } +}