diff --git a/vnt/Cargo.toml b/vnt/Cargo.toml index 91ebcc6..2af95d8 100644 --- a/vnt/Cargo.toml +++ b/vnt/Cargo.toml @@ -11,7 +11,7 @@ bytes = "1.3.0" log = "0.4.17" libc = "0.2.137" crossbeam-utils = "0.8" -crossbeam-skiplist = "0.1.1" +dashmap = "5.5.1" parking_lot = "0.12.1" rand = "0.8.5" sha2 = { version = "0.10.6", features = ["oid"] } diff --git a/vnt/src/channel/channel.rs b/vnt/src/channel/channel.rs index f3486dd..60dd83f 100644 --- a/vnt/src/channel/channel.rs +++ b/vnt/src/channel/channel.rs @@ -2,9 +2,8 @@ use std::io; use std::net::{Ipv4Addr, SocketAddr}; use std::sync::Arc; use std::time::{Duration, Instant}; -use crossbeam_skiplist::SkipMap; use crossbeam_utils::atomic::AtomicCell; -use parking_lot::Mutex; +use dashmap::DashMap; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpStream, UdpSocket}; use tokio::net::tcp::OwnedReadHalf; @@ -16,16 +15,15 @@ use crate::handle::CurrentDeviceInfo; use crate::handle::recv_handler::ChannelDataHandler; pub struct ContextInner { - pub(crate) lock: Mutex<()>, //udp用于打洞、服务端通信(可选) pub(crate) main_channel: Arc, //在udp的基础上,可以选择使用tcp和服务端通信 pub(crate) main_tcp_channel: Option>>, - pub(crate) route_table: SkipMap>, - pub(crate) route_table_time: SkipMap<(RouteKey, Ipv4Addr), AtomicCell>, + pub(crate) route_table: DashMap>, + pub(crate) route_table_time: DashMap<(RouteKey, Ipv4Addr), AtomicCell>, pub(crate) status_receiver: Receiver, pub(crate) status_sender: Sender, - pub(crate) udp_map: SkipMap>, + pub(crate) udp_map: DashMap>, pub(crate) channel_num: usize, current_device: Arc>, } @@ -41,14 +39,13 @@ impl Context { let channel_num = 1; let (status_sender, status_receiver) = channel(Status::Cone); let inner = Arc::new(ContextInner { - lock: Mutex::new(()), main_channel, main_tcp_channel, - route_table: SkipMap::new(), - route_table_time: SkipMap::new(), + route_table: DashMap::with_capacity(16), + route_table_time: DashMap::with_capacity(16), status_receiver, status_sender, - udp_map: SkipMap::new(), + udp_map: DashMap::new(), channel_num, current_device, }); @@ -117,8 +114,10 @@ impl Context { } pub(crate) async fn send_all(&self, buf: &[u8], addr: SocketAddr) -> io::Result<()> { - for udp in self.inner.udp_map.iter() { - udp.value().send_to(buf, addr).await?; + for udp_ref in self.inner.udp_map.iter() { + let udp = udp_ref.clone(); + drop(udp_ref); + udp.send_to(buf, addr).await?; } Ok(()) } @@ -139,8 +138,10 @@ impl Context { } } - if let Some(udp) = self.inner.udp_map.get(&route.index) { - return udp.value().send_to(buf, route.addr).await; + if let Some(udp_ref) = self.inner.udp_map.get(&route.index) { + let udp = udp_ref.value().clone(); + drop(udp_ref); + return udp.send_to(buf, route.addr).await; } } Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) @@ -171,8 +172,10 @@ impl Context { }; } } - if let Some(udp) = self.inner.udp_map.get(&route_key.index) { - return udp.value().send_to(buf, route_key.addr).await; + if let Some(udp_ref) = self.inner.udp_map.get(&route_key.index) { + let udp = udp_ref.value().clone(); + drop(udp_ref); + return udp.send_to(buf, route_key.addr).await; } Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) } @@ -201,12 +204,7 @@ impl Context { } fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) { let key = route.route_key(); - let guard = self.inner.lock.lock(); - let mut list = if let Some(entry) = self.inner.route_table.get(&id) { - entry.value().clone() - } else { - Vec::with_capacity(4) - }; + let mut list = self.inner.route_table.entry(id).or_insert_with(||Vec::with_capacity(4)); let mut exist = false; for x in list.iter_mut() { if x.metric < route.metric { @@ -237,9 +235,7 @@ impl Context { list.truncate(max_len); } } - self.inner.route_table.insert(id, list); self.inner.route_table_time.insert((key, id), AtomicCell::new(Instant::now())); - drop(guard); } pub fn route(&self, id: &Ipv4Addr) -> Option> { if let Some(v) = self.inner.route_table.get(id) { @@ -297,16 +293,13 @@ impl Context { v } pub fn remove_route_all(&self, id: &Ipv4Addr) { - let guard = self.inner.lock.lock(); - if let Some(v) = self.inner.route_table.remove(id) { - for x in v.value() { + if let Some((_,routes)) = self.inner.route_table.remove(id) { + for x in routes { self.inner.route_table_time.remove(&(x.route_key(), *id)); } } - drop(guard); } pub fn remove_route(&self, id: &Ipv4Addr, route_key: RouteKey) { - let guard = self.inner.lock.lock(); if let Some(v) = self.inner.route_table.get(id) { let mut routes = v.value().clone(); drop(v); @@ -314,7 +307,6 @@ impl Context { self.inner.route_table.insert(*id, routes); } self.inner.route_table_time.remove(&(route_key, *id)); - drop(guard); } pub fn update_read_time(&self, id: &Ipv4Addr, route_key: &RouteKey) { if let Some(time) = self.inner.route_table_time.get(&(*route_key, *id)) { diff --git a/vnt/src/channel/idle.rs b/vnt/src/channel/idle.rs index 4358331..06c95b7 100644 --- a/vnt/src/channel/idle.rs +++ b/vnt/src/channel/idle.rs @@ -29,7 +29,6 @@ impl Idle { for entry in self.context.inner.route_table_time.iter() { let last_read = entry.value().load().elapsed(); if last_read >= self.read_idle { - entry.remove(); return Ok((entry.key().1.clone(), entry.key().0.clone())); } else { if max < last_read { diff --git a/vnt/src/core/mod.rs b/vnt/src/core/mod.rs index a97c010..cdbc185 100644 --- a/vnt/src/core/mod.rs +++ b/vnt/src/core/mod.rs @@ -3,8 +3,8 @@ use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; use std::sync::Arc; use std::time::Duration; -use crossbeam_skiplist::SkipMap; use crossbeam_utils::atomic::AtomicCell; +use dashmap::DashMap; use parking_lot::Mutex; use rand::Rng; use tokio::net::{TcpStream, UdpSocket}; @@ -48,7 +48,7 @@ pub struct Vnt { device_list: Arc)>>, nat_test: NatTest, connect_status: Arc>, - peer_nat_info_map: Arc>, + peer_nat_info_map: Arc>, } pub struct VntUtil { @@ -208,7 +208,7 @@ impl VntUtil { config.server_address, config.token.clone(), config.device_id.clone(), config.name.clone(),config.password.is_some())); let device_list: Arc)>> = Arc::new(Mutex::new((response.epoch, response.device_info_list))); - let peer_nat_info_map: Arc> = Arc::new(SkipMap::new()); + let peer_nat_info_map: Arc> = Arc::new(DashMap::new()); let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected)); diff --git a/vnt/src/handle/recv_handler.rs b/vnt/src/handle/recv_handler.rs index 82725e8..89c70bd 100644 --- a/vnt/src/handle/recv_handler.rs +++ b/vnt/src/handle/recv_handler.rs @@ -1,8 +1,8 @@ use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; use std::sync::Arc; -use crossbeam_skiplist::SkipMap; use crossbeam_utils::atomic::AtomicCell; +use dashmap::DashMap; use parking_lot::Mutex; use protobuf::Message; use tokio::sync::mpsc::Sender; @@ -41,7 +41,7 @@ pub struct ChannelDataHandler { igmp_server: Option, device_writer: DeviceWriter, connect_status: Arc>, - peer_nat_info_map: Arc>, + peer_nat_info_map: Arc>, ip_proxy_map: Option, out_external_route: AllowExternalRoute, cone_sender: Sender<(Ipv4Addr, NatInfo)>, @@ -61,7 +61,7 @@ impl ChannelDataHandler { igmp_server: Option, device_writer: DeviceWriter, connect_status: Arc>, - peer_nat_info_map: Arc>, + peer_nat_info_map: Arc>, ip_proxy_map: Option, out_external_route: AllowExternalRoute, cone_sender: Sender<(Ipv4Addr, NatInfo)>, @@ -196,9 +196,7 @@ impl ChannelDataHandler { ipv4.update_checksum(); let key = SocketAddrV4::new(source, source_port); //https://github.com/crossbeam-rs/crossbeam/issues/1023 - if !ip_proxy_map.tcp_proxy_map.contains_key(&key){ - ip_proxy_map.tcp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); - } + ip_proxy_map.tcp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); } ipv4::protocol::Protocol::Udp => { let dest_ip = ipv4.destination_ip(); @@ -211,9 +209,7 @@ impl ChannelDataHandler { ipv4.set_destination_ip(destination); ipv4.update_checksum(); let key = SocketAddrV4::new(source, source_port); - if !ip_proxy_map.udp_proxy_map.contains_key(&key){ - ip_proxy_map.udp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); - } + ip_proxy_map.udp_proxy_map.insert(key, SocketAddrV4::new(dest_ip, dest_port)); } ipv4::protocol::Protocol::Icmp => { let dest_ip = ipv4.destination_ip(); diff --git a/vnt/src/igmp_server/mod.rs b/vnt/src/igmp_server/mod.rs index 348e568..f6bcad5 100644 --- a/vnt/src/igmp_server/mod.rs +++ b/vnt/src/igmp_server/mod.rs @@ -2,7 +2,7 @@ use std::collections::{HashMap, HashSet}; use std::net::Ipv4Addr; use std::sync::Arc; use std::time::{Duration, Instant}; -use crossbeam_skiplist::SkipMap; +use dashmap::DashMap; use parking_lot::RwLock; use packet::igmp::igmp_v2::IgmpV2Packet; use packet::igmp::igmp_v3::{IgmpV3QueryPacket, IgmpV3RecordType, IgmpV3ReportPacket}; @@ -47,12 +47,12 @@ impl Multicast { #[derive(Clone)] pub struct IgmpServer { - multicast: Arc>>>, + multicast: Arc>>>, } impl IgmpServer { pub fn new(device_writer: DeviceWriter) -> Self { - let multicast: Arc>>> = Arc::new(SkipMap::new()); + let multicast: Arc>>> = Arc::new(DashMap::new()); std::thread::spawn(move || { //预留以太网帧头和ip头 let mut buf = [0; 14 + 24 + 12]; @@ -124,10 +124,12 @@ impl IgmpServer { if !multicast_addr.is_multicast() { return Ok(()); } - let multi = self.multicast.get_or_insert_with(multicast_addr, || { - Arc::new(RwLock::new(Multicast::new())) - }); - let mut guard = multi.value().write(); + let multi = { + self.multicast.entry(multicast_addr).or_insert_with(|| { + Arc::new(RwLock::new(Multicast::new())) + }).value().clone() + }; + let mut guard = multi.write(); guard.members.insert(source, Instant::now()); } IgmpType::LeaveV2 => { @@ -151,10 +153,10 @@ impl IgmpServer { if !multicast_addr.is_multicast() { return Ok(()); } - let multi = self.multicast.get_or_insert_with(multicast_addr, || { + let multi = self.multicast.entry(multicast_addr).or_insert_with(|| { Arc::new(RwLock::new(Multicast::new())) - }); - let mut guard = multi.value().write(); + }).value().clone(); + let mut guard = multi.write(); match group_record.record_type() { IgmpV3RecordType::ModeIsInclude | IgmpV3RecordType::ChangeToIncludeMode => { diff --git a/vnt/src/ip_proxy/icmp_proxy.rs b/vnt/src/ip_proxy/icmp_proxy.rs index cf32767..2d3c19e 100644 --- a/vnt/src/ip_proxy/icmp_proxy.rs +++ b/vnt/src/ip_proxy/icmp_proxy.rs @@ -3,8 +3,8 @@ use std::mem::MaybeUninit; use std::net::{IpAddr, Ipv4Addr, SocketAddrV4}; use std::sync::Arc; use crossbeam_utils::atomic::AtomicCell; +use dashmap::DashMap; -use crossbeam_skiplist::SkipMap; use socket2::{Domain, SockAddr, Socket, Type}; use packet::icmp::icmp; @@ -19,14 +19,14 @@ use crate::protocol::body::ENCRYPTION_RESERVED; pub struct IcmpProxy { icmp_socket: Arc, // 对端-> 真实来源 - icmp_proxy_map: Arc>, + icmp_proxy_map: Arc>, sender: ChannelSender, current_device: Arc>, client_cipher: Cipher, } impl IcmpProxy { - pub fn new(addr: SocketAddrV4, icmp_proxy_map: Arc>, + pub fn new(addr: SocketAddrV4, icmp_proxy_map: Arc>, sender: ChannelSender, current_device: Arc>, client_cipher: Cipher) -> io::Result { let icmp_socket = Arc::new(Socket::new(Domain::IPV4, Type::RAW, Some(socket2::Protocol::ICMPV4))?); icmp_socket.bind(&SockAddr::from(addr))?; @@ -60,6 +60,7 @@ impl IcmpProxy { if let Some(entry) = self.icmp_proxy_map.get(&(peer_ip, id, seq)) { //将数据发送到真实的来源 let dest_ip = *entry.value(); + drop(entry); ipv4_packet.set_destination_ip(dest_ip); ipv4_packet.update_checksum(); let current_device = self.current_device.load(); diff --git a/vnt/src/ip_proxy/mod.rs b/vnt/src/ip_proxy/mod.rs index 8497272..933f682 100644 --- a/vnt/src/ip_proxy/mod.rs +++ b/vnt/src/ip_proxy/mod.rs @@ -2,7 +2,7 @@ use std::{io, thread}; use std::net::{Ipv4Addr, SocketAddrV4}; use std::sync::Arc; use crossbeam_utils::atomic::AtomicCell; -use crossbeam_skiplist::SkipMap; +use dashmap::DashMap; use socket2::{SockAddr, Socket}; use tokio::net::{TcpListener, UdpSocket}; use crate::channel::sender::ChannelSender; @@ -28,10 +28,10 @@ pub struct IpProxyMap { pub(crate) tcp_proxy_port: u16, pub(crate) udp_proxy_port: u16, //真实源地址 -> 目的地址 - pub(crate) tcp_proxy_map: Arc>, - pub(crate) udp_proxy_map: Arc>, + pub(crate) tcp_proxy_map: Arc>, + pub(crate) udp_proxy_map: Arc>, // icmp用Identifier来区分,没有Identifier的一律不转发 - pub(crate) icmp_proxy_map: Arc>, + pub(crate) icmp_proxy_map: Arc>, icmp_socket: Arc, } @@ -42,9 +42,9 @@ impl IpProxyMap { } pub async fn init_proxy(sender: ChannelSender, current_device: Arc>, client_cipher: Cipher,) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> { - let tcp_proxy_map: Arc> = Arc::new(SkipMap::new()); - let udp_proxy_map: Arc> = Arc::new(SkipMap::new()); - let icmp_proxy_map: Arc> = Arc::new(SkipMap::new()); + let tcp_proxy_map: Arc> = Arc::new(DashMap::new()); + let udp_proxy_map: Arc> = Arc::new(DashMap::new()); + let icmp_proxy_map: Arc> = Arc::new(DashMap::new()); let tcp_listener = TcpListener::bind("0.0.0.0:0").await?; let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; let tcp_proxy_port = tcp_listener.local_addr()?.port(); diff --git a/vnt/src/ip_proxy/tcp_proxy.rs b/vnt/src/ip_proxy/tcp_proxy.rs index 4a4c640..42df18d 100644 --- a/vnt/src/ip_proxy/tcp_proxy.rs +++ b/vnt/src/ip_proxy/tcp_proxy.rs @@ -1,17 +1,17 @@ use std::io; use std::net::{SocketAddr, SocketAddrV4}; use std::sync::Arc; +use dashmap::DashMap; -use crossbeam_skiplist::SkipMap; use tokio::net::{TcpListener, TcpStream}; pub struct TcpProxy { tcp_listener: TcpListener, - tcp_proxy_map: Arc>, + tcp_proxy_map: Arc>, } impl TcpProxy { - pub fn new(tcp_listener: TcpListener, tcp_proxy_map: Arc>) -> Self { + pub fn new(tcp_listener: TcpListener, tcp_proxy_map: Arc>) -> Self { Self { tcp_listener, tcp_proxy_map, @@ -27,6 +27,7 @@ impl TcpProxy { SocketAddr::V4(sender_addr) => { if let Some(entry) = tcp_proxy_map.get(&sender_addr) { let dest_addr = *entry.value(); + drop(entry); let peer_tcp_stream = match TcpStream::connect(dest_addr).await { Ok(peer_tcp_stream) => { peer_tcp_stream } Err(e) => { diff --git a/vnt/src/ip_proxy/udp_proxy.rs b/vnt/src/ip_proxy/udp_proxy.rs index 7cb7d01..caa310a 100644 --- a/vnt/src/ip_proxy/udp_proxy.rs +++ b/vnt/src/ip_proxy/udp_proxy.rs @@ -2,17 +2,17 @@ use std::io; use std::net::{SocketAddr, SocketAddrV4}; use std::sync::Arc; use std::time::Duration; -use crossbeam_skiplist::SkipMap; +use dashmap::DashMap; use tokio::net::UdpSocket; /// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址 pub struct UdpProxy { udp_socket: Arc, - map: Arc>, + map: Arc>, } impl UdpProxy { - pub fn new(udp_socket: UdpSocket, map: Arc>) -> Self { + pub fn new(udp_socket: UdpSocket, map: Arc>) -> Self { let udp_socket = Arc::new(udp_socket); Self { udp_socket, @@ -23,7 +23,7 @@ impl UdpProxy { let map = self.map; let udp_socket = self.udp_socket; let mut buf = [0u8; 65536]; - let inner_map: Arc>> = Arc::new(SkipMap::new()); + let inner_map: Arc>> = Arc::new(DashMap::new()); loop { match udp_socket.recv_from(&mut buf).await { @@ -48,11 +48,14 @@ impl UdpProxy { } } -async fn start0(buf: &[u8], sender_addr: SocketAddrV4, inner_map: &Arc>>, map: &Arc>, udp_socket: &Arc) -> io::Result<()> { +async fn start0(buf: &[u8], sender_addr: SocketAddrV4, inner_map: &Arc>>, map: &Arc>, udp_socket: &Arc) -> io::Result<()> { if let Some(entry) = inner_map.get(&sender_addr) { - entry.value().send(buf).await?; + let udp = entry.value().clone(); + drop(entry); + udp.send(buf).await?; } else if let Some(entry) = map.get(&sender_addr) { let dest_addr = *entry.value(); + drop(entry); let peer_udp_socket = UdpSocket::bind("0.0.0.0:0").await?; peer_udp_socket.connect(dest_addr).await?; peer_udp_socket.send(buf).await?;