diff --git a/vnt/src/ip_proxy/icmp_proxy.rs b/vnt/src/ip_proxy/icmp_proxy.rs index 769f255..f2b1845 100644 --- a/vnt/src/ip_proxy/icmp_proxy.rs +++ b/vnt/src/ip_proxy/icmp_proxy.rs @@ -1,137 +1,206 @@ -use std::io; -use std::mem::MaybeUninit; -use std::net::{IpAddr, Ipv4Addr, SocketAddrV4}; +use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}; use std::sync::Arc; +use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; -use dashmap::DashMap; -use socket2::{Domain, SockAddr, Socket, Type}; +use mio::net::UdpSocket; +use mio::{Events, Interest, Poll, Token, Waker}; +use parking_lot::Mutex; use packet::icmp::icmp; use packet::icmp::icmp::HeaderOther; use packet::ip::ipv4::packet::IpV4Packet; -use crate::channel::sender::ChannelSender; +use crate::channel::context::Context; use crate::cipher::Cipher; use crate::handle::CurrentDeviceInfo; -use crate::ip_proxy::{send, ProxyHandler}; - +use crate::ip_proxy::ProxyHandler; +use crate::protocol; +use crate::protocol::{NetPacket, Version, MAX_TTL}; +use crate::util::StopManager; +#[derive(Clone)] pub struct IcmpProxy { - icmp_socket: Arc, + icmp_socket: Arc, // 对端-> 真实来源 - icmp_proxy_map: Arc>, - sender: ChannelSender, - current_device: Arc>, - client_cipher: Cipher, + nat_map: Arc>>, } impl IcmpProxy { pub fn new( - addr: SocketAddrV4, - icmp_proxy_map: Arc>, - sender: ChannelSender, + context: Context, + stop_manager: StopManager, current_device: Arc>, client_cipher: Cipher, - ) -> io::Result { - let icmp_socket = Arc::new(Socket::new( - Domain::IPV4, - Type::RAW, + ) -> io::Result { + let icmp_socket = socket2::Socket::new( + socket2::Domain::IPV4, + socket2::Type::RAW, Some(socket2::Protocol::ICMPV4), - )?); - icmp_socket.bind(&SockAddr::from(addr))?; - Ok(IcmpProxy { - icmp_socket, - icmp_proxy_map, - sender, - current_device, - client_cipher, + )?; + let addr: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); + icmp_socket.bind(&socket2::SockAddr::from(addr))?; + icmp_socket.set_nonblocking(true)?; + let std_socket: std::net::UdpSocket = icmp_socket.into(); + let mio_icmp_socket = UdpSocket::from_std(std_socket.try_clone()?); + let nat_map: Arc>> = + Arc::new(Mutex::new(HashMap::with_capacity(16))); + { + let nat_map = nat_map.clone(); + thread::spawn(move || { + if let Err(e) = icmp_proxy( + mio_icmp_socket, + nat_map, + context, + stop_manager, + current_device, + client_cipher, + ) { + log::warn!("icmp_proxy:{:?}", e); + } + }); + } + Ok(Self { + icmp_socket: Arc::new(std_socket), + nat_map, }) } - pub fn icmp_handler(&self) -> IcmpHandler { - IcmpHandler(self.icmp_socket.clone(), self.icmp_proxy_map.clone()) - } - pub fn start(self) { - let mut buf = [0u8; 4096]; - let data: &mut [MaybeUninit] = unsafe { std::mem::transmute(&mut buf[12..]) }; +} - loop { - match self.recv(data) { - Ok((len, peer_ip)) => { - match peer_ip { - IpAddr::V4(peer_ip) => { - match IpV4Packet::new(&mut buf[12..12 + len]) { - Ok(mut ipv4_packet) => { - match icmp::IcmpPacket::new(ipv4_packet.payload()) { - Ok(icmp_packet) => { - match icmp_packet.header_other() { - HeaderOther::Identifier(id, seq) => { - 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(); - send( - &mut buf, - len, - dest_ip, - &self.sender, - &self.current_device, - &self.client_cipher, - ); - } - } - _ => { - continue; - } - } - } - Err(_) => {} - }; - } - Err(_) => {} +const SERVER_VAL: usize = 0; +const SERVER: Token = Token(SERVER_VAL); +const NOTIFY_VAL: usize = 1; +const NOTIFY: Token = Token(NOTIFY_VAL); + +fn icmp_proxy( + mut icmp_socket: UdpSocket, + // 对端-> 真实来源 + nat_map: Arc>>, + context: Context, + stop_manager: StopManager, + current_device: Arc>, + client_cipher: Cipher, +) -> io::Result<()> { + let mut poll = Poll::new()?; + poll.registry() + .register(&mut icmp_socket, SERVER, Interest::READABLE)?; + let mut events = Events::with_capacity(32); + let stop = Waker::new(poll.registry(), NOTIFY)?; + let _worker = stop_manager.add_listener("icmp_proxy".into(), move || { + if let Err(e) = stop.wake() { + log::warn!("stop icmp_proxy:{:?}", e); + } + })?; + let mut buf = [0u8; 65535 - 20 - 8]; + loop { + poll.poll(&mut events, None)?; + + for event in events.iter() { + match event.token() { + SERVER => loop { + let (len, addr) = match icmp_socket.recv_from(&mut buf[12..]) { + Ok(rs) => rs, + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; } + log::warn!("icmp_socket {:?}", e); + break; } - IpAddr::V6(_) => {} + }; + if let IpAddr::V4(peer_ip) = addr.ip() { + recv_handle( + &mut buf, + 12 + len, + peer_ip, + &nat_map, + &context, + ¤t_device, + &client_cipher, + ); } + }, + NOTIFY => { + return Ok(()); } - Err(e) => { - log::warn!("icmp代理异常:{:?}", e); - } + _ => {} } } } - fn recv(&self, buf: &mut [MaybeUninit]) -> io::Result<(usize, IpAddr)> { - let (size, addr) = self.icmp_socket.recv_from(buf)?; - let addr = match addr.as_socket() { - None => IpAddr::V4(Ipv4Addr::UNSPECIFIED), - Some(add) => add.ip(), - }; - Ok((size, addr)) +} + +fn recv_handle( + buf: &mut [u8], + data_len: usize, + peer_ip: Ipv4Addr, + nat_map: &Mutex>, + context: &Context, + current_device: &AtomicCell, + client_cipher: &Cipher, +) { + match IpV4Packet::new(&mut buf[12..data_len]) { + Ok(mut ipv4_packet) => match icmp::IcmpPacket::new(ipv4_packet.payload()) { + Ok(icmp_packet) => match icmp_packet.header_other() { + HeaderOther::Identifier(id, seq) => { + if let Some(dest_ip) = nat_map.lock().remove(&(peer_ip, id, seq)) { + ipv4_packet.set_destination_ip(dest_ip); + ipv4_packet.update_checksum(); + + let current_device = current_device.load(); + let virtual_ip = current_device.virtual_ip(); + + let mut net_packet = NetPacket::new0(data_len, buf).unwrap(); + net_packet.set_version(Version::V1); + net_packet.set_protocol(protocol::Protocol::IpTurn); + net_packet.set_transport_protocol( + protocol::ip_turn_packet::Protocol::Ipv4.into(), + ); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_source(virtual_ip); + net_packet.set_destination(dest_ip); + if let Err(e) = client_cipher.encrypt_ipv4(&mut net_packet) { + log::warn!("加密失败:{}", e); + return; + } + if context.send_by_id(net_packet.buffer(), &dest_ip).is_err() { + let connect_server = current_device.connect_server; + if let Err(e) = + context.send_default(net_packet.buffer(), connect_server) + { + log::warn!("发送到目标失败:{},{}", e, connect_server); + } + } + } + } + _ => {} + }, + Err(_) => {} + }, + Err(_) => {} } } -/// icmp用Identifier来区分,没有Identifier的一律不转发 -#[derive(Clone)] -pub struct IcmpHandler(Arc, Arc>); -impl ProxyHandler for IcmpHandler { +/// icmp用Identifier来区分,没有Identifier的一律不转发 +impl ProxyHandler for IcmpProxy { fn recv_handle( &self, ipv4: &mut IpV4Packet<&mut [u8]>, source: Ipv4Addr, destination: Ipv4Addr, ) -> io::Result { + if ipv4.offset() != 0 || ipv4.flags() & 1 == 1 { + // ip分片的直接丢弃 + return Ok(true); + } let dest_ip = ipv4.destination_ip(); //转发到代理目标地址 let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?; match icmp_packet.header_other() { HeaderOther::Identifier(id, seq) => { - self.1.insert((dest_ip, id, seq), source); - self.0.send_to( + self.nat_map.lock().insert((dest_ip, id, seq), source); + self.icmp_socket.send_to( ipv4.payload(), - &SockAddr::from(SocketAddrV4::new(dest_ip, 0)), + SocketAddr::from(SocketAddrV4::new(dest_ip, 0)), )?; } _ => { diff --git a/vnt/src/ip_proxy/mod.rs b/vnt/src/ip_proxy/mod.rs index 9e70552..38cba05 100644 --- a/vnt/src/ip_proxy/mod.rs +++ b/vnt/src/ip_proxy/mod.rs @@ -1,35 +1,24 @@ +use std::io; use std::net::Ipv4Addr; -use std::net::SocketAddrV4; use std::sync::Arc; -use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; -use dashmap::DashMap; -use packet::ip::ipv4; -use tokio::net::UdpSocket; +use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; -use crate::channel::sender::ChannelSender; +use crate::channel::context::Context; use crate::cipher::Cipher; use crate::handle::CurrentDeviceInfo; -#[cfg(not(target_os = "android"))] -use crate::ip_proxy::icmp_proxy::IcmpHandler; -use crate::ip_proxy::tcp_proxy::{TcpHandler, TcpProxy}; -use crate::ip_proxy::udp_proxy::{UdpHandler, UdpProxy}; -use crate::protocol; -use crate::protocol::{NetPacket, Version, MAX_TTL}; +use crate::ip_proxy::icmp_proxy::IcmpProxy; +use crate::ip_proxy::tcp_proxy::TcpProxy; +use crate::ip_proxy::udp_proxy::UdpProxy; +use crate::util::{Scheduler, StopManager}; -#[cfg(not(target_os = "android"))] pub mod icmp_proxy; pub mod tcp_proxy; pub mod udp_proxy; -pub trait DashMapNew { - fn new0() -> Self; - fn new_cap(capacity: usize) -> Self; -} - pub trait ProxyHandler { fn recv_handle( &self, @@ -40,139 +29,29 @@ pub trait ProxyHandler { fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()>; } -impl<'a, K: 'a + Eq + std::hash::Hash, V: 'a> DashMapNew for DashMap { - fn new0() -> Self { - Self::new_cap(0) - } - - fn new_cap(capacity: usize) -> Self { - let shard_amount = (thread::available_parallelism().map_or(4, |v| { - // https://github.com/rust-lang/rust/issues/115868 - let n: usize = v.get() * 4; - if n == 0 { - log::warn!("available_parallelism=0"); - println!("warn available_parallelism=0"); - } - if n < 4 { - return 4; - } - n - })) - .next_power_of_two(); - DashMap::with_capacity_and_shard_amount(capacity, shard_amount) - } -} - -#[derive(Eq, PartialEq, Ord, PartialOrd, Copy, Clone, Debug)] -pub enum Protocol { - Icmp, - Tcp, - Udp, -} - #[derive(Clone)] pub struct IpProxyMap { - #[cfg(not(target_os = "android"))] - pub(crate) icmp_handler: IcmpHandler, - pub(crate) tcp_handler: TcpHandler, - pub(crate) udp_handler: UdpHandler, + icmp_proxy: IcmpProxy, + tcp_proxy: TcpProxy, + udp_proxy: UdpProxy, } -#[cfg(not(target_os = "android"))] -pub async fn init_proxy( - sender: ChannelSender, +pub fn init_proxy( + context: Context, + scheduler: Scheduler, + stop_manager: StopManager, current_device: Arc>, client_cipher: Cipher, -) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> { - let tcp_proxy_map: Arc> = Arc::new(DashMap::new0()); - let udp_proxy_map: Arc> = Arc::new(DashMap::new0()); +) -> io::Result { + let icmp_proxy = IcmpProxy::new(context, stop_manager.clone(), current_device, client_cipher)?; + let tcp_proxy = TcpProxy::new(stop_manager.clone())?; + let udp_proxy = UdpProxy::new(scheduler, stop_manager)?; - let icmp_proxy_map: Arc> = Arc::new(DashMap::new0()); - let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; - let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); - let tcp_proxy = TcpProxy::new(addr, tcp_proxy_map.clone()).await?; - let tcp_handler = tcp_proxy.tcp_handler(); - let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone())?; - let udp_handler = udp_proxy.udp_handler(); - - let icmp_handler = { - let icmp_proxy = icmp_proxy::IcmpProxy::new( - addr, - icmp_proxy_map.clone(), - sender.clone(), - current_device.clone(), - client_cipher.clone(), - )?; - let icmp_handler = icmp_proxy.icmp_handler(); - thread::spawn(move || { - icmp_proxy.start(); - }); - icmp_handler - }; - - Ok(( + Ok(IpProxyMap { + icmp_proxy, tcp_proxy, udp_proxy, - IpProxyMap { - tcp_handler, - udp_handler, - icmp_handler, - }, - )) -} - -#[cfg(target_os = "android")] -pub async fn init_proxy() -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> { - let tcp_proxy_map: Arc> = Arc::new(DashMap::new0()); - let udp_proxy_map: Arc> = Arc::new(DashMap::new0()); - let addr = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); - let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; - let tcp_proxy = TcpProxy::new(addr, tcp_proxy_map.clone()).await?; - let tcp_handler = tcp_proxy.tcp_handler(); - let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone())?; - let udp_handler = udp_proxy.udp_handler(); - - Ok(( - tcp_proxy, - udp_proxy, - IpProxyMap { - tcp_handler, - udp_handler, - }, - )) -} - -pub fn send( - buf: &mut [u8], - data_len: usize, - dest_ip: Ipv4Addr, - sender: &ChannelSender, - current_device: &AtomicCell, - client_cipher: &Cipher, -) { - let current_device = current_device.load(); - let virtual_ip = current_device.virtual_ip(); - - let mut net_packet = NetPacket::new0(12 + data_len, buf).unwrap(); - net_packet.set_version(Version::V1); - net_packet.set_protocol(protocol::Protocol::IpTurn); - net_packet.set_transport_protocol(protocol::ip_turn_packet::Protocol::Ipv4.into()); - net_packet.first_set_ttl(MAX_TTL); - net_packet.set_source(virtual_ip); - net_packet.set_destination(dest_ip); - if let Err(e) = client_cipher.encrypt_ipv4(&mut net_packet) { - log::warn!("加密失败:{}", e); - return; - } - if sender - .try_send_by_id(net_packet.buffer(), &dest_ip) - .is_err() - { - let connect_server = current_device.connect_server; - if let Err(e) = sender.send_main(net_packet.buffer(), connect_server) { - log::warn!("发送到目标失败:{},{}", e, connect_server); - } - } + }) } impl ProxyHandler for IpProxyMap { @@ -183,15 +62,10 @@ impl ProxyHandler for IpProxyMap { destination: Ipv4Addr, ) -> io::Result { match ipv4.protocol() { - ipv4::protocol::Protocol::Tcp => { - self.tcp_handler.recv_handle(ipv4, source, destination) - } - ipv4::protocol::Protocol::Udp => { - self.udp_handler.recv_handle(ipv4, source, destination) - } - #[cfg(not(target_os = "android"))] + ipv4::protocol::Protocol::Tcp => self.tcp_proxy.recv_handle(ipv4, source, destination), + ipv4::protocol::Protocol::Udp => self.udp_proxy.recv_handle(ipv4, source, destination), ipv4::protocol::Protocol::Icmp => { - self.icmp_handler.recv_handle(ipv4, source, destination) + self.icmp_proxy.recv_handle(ipv4, source, destination) } _ => { log::warn!( @@ -208,8 +82,9 @@ impl ProxyHandler for IpProxyMap { fn send_handle(&self, ipv4: &mut IpV4Packet<&mut [u8]>) -> io::Result<()> { match ipv4.protocol() { - ipv4::protocol::Protocol::Tcp => self.tcp_handler.send_handle(ipv4), - ipv4::protocol::Protocol::Udp => self.udp_handler.send_handle(ipv4), + ipv4::protocol::Protocol::Tcp => self.tcp_proxy.send_handle(ipv4), + ipv4::protocol::Protocol::Udp => self.udp_proxy.send_handle(ipv4), + ipv4::protocol::Protocol::Icmp => self.icmp_proxy.send_handle(ipv4), _ => Ok(()), } } diff --git a/vnt/src/ip_proxy/tcp_proxy.rs b/vnt/src/ip_proxy/tcp_proxy.rs index 1d73bce..598a014 100644 --- a/vnt/src/ip_proxy/tcp_proxy.rs +++ b/vnt/src/ip_proxy/tcp_proxy.rs @@ -1,141 +1,53 @@ -use crate::ip_proxy::ProxyHandler; -use crossbeam_utils::atomic::AtomicCell; -use dashmap::DashMap; +use std::io::{Read, Write}; +use std::net::{Ipv4Addr, Shutdown, SocketAddrV4}; +#[cfg(unix)] +use std::os::fd::AsRawFd; +#[cfg(windows)] +use std::os::windows::io::AsRawSocket; +use std::sync::Arc; +use std::{collections::HashMap, io, net::SocketAddr, thread}; + +use bytes::{BufMut, BytesMut}; +use mio::net::TcpStream; +use mio::{net::TcpListener, Events, Interest, Poll, Registry, Token, Waker}; +use parking_lot::Mutex; + use packet::ip::ipv4::packet::IpV4Packet; use packet::tcp::tcp::TcpPacket; -use std::io; -use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; -use std::sync::Arc; -use std::time::{Duration, Instant}; -use tokio::io::AsyncReadExt; -use tokio::io::AsyncWriteExt; -use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; -use tokio::net::{TcpListener, TcpStream}; +use crate::ip_proxy::ProxyHandler; +use crate::util::StopManager; + +const SERVER_VAL: usize = 0; +const SERVER: Token = Token(SERVER_VAL); +const NOTIFY_VAL: usize = 1; +const NOTIFY: Token = Token(NOTIFY_VAL); + +#[derive(Clone)] pub struct TcpProxy { - tcp_proxy_port: u16, - tcp_listener: TcpListener, - //真实源地址 -> 目的地址 - tcp_proxy_map: Arc>, + port: u16, + nat_map: Arc>>, } impl TcpProxy { - pub async fn new( - addr: SocketAddrV4, - tcp_proxy_map: Arc>, - ) -> io::Result { - let tcp_listener = TcpListener::bind(addr).await?; - Ok(Self { - tcp_proxy_port: tcp_listener.local_addr()?.port(), - tcp_listener, - tcp_proxy_map, - }) - } - pub fn tcp_handler(&self) -> TcpHandler { - TcpHandler(self.tcp_proxy_port, self.tcp_proxy_map.clone()) - } - pub async fn start(self) { - let tcp_listener = self.tcp_listener; - let tcp_proxy_map = self.tcp_proxy_map; - loop { - match tcp_listener.accept().await { - Ok((tcp_stream, sender_addr)) => match sender_addr { - SocketAddr::V4(sender_addr) => { - if let Some(entry) = tcp_proxy_map.get(&sender_addr) { - let dest_addr = *entry.value(); - drop(entry); - - tokio::spawn(async move { - let peer_tcp_stream = match tokio::time::timeout( - Duration::from_secs(5), - TcpStream::connect(dest_addr), - ) - .await - { - Ok(peer_tcp_stream) => match peer_tcp_stream { - Ok(peer_tcp_stream) => peer_tcp_stream, - Err(e) => { - log::warn!( - "tcp代理异常:{:?},来源:{},目标:{}", - e, - sender_addr, - dest_addr - ); - return; - } - }, - Err(e) => { - log::warn!( - "tcp代理异常:{:?},来源:{},目标:{}", - e, - sender_addr, - dest_addr - ); - return; - } - }; - if let Err(e) = proxy(tcp_stream, peer_tcp_stream).await { - log::warn!("{}->{},{}", sender_addr, dest_addr, e); - } - }); - } else { - log::warn!("tcp代理异常: 来源:{},未找到目标", sender_addr); - } - } - SocketAddr::V6(_) => {} - }, - Err(e) => { - log::warn!("tcp代理监听:{:?}", e); + pub fn new(stop_manager: StopManager) -> io::Result { + let nat_map: Arc>> = + Arc::new(Mutex::new(HashMap::with_capacity(16))); + let tcp_listener = TcpListener::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?; + let port = tcp_listener.local_addr()?.port(); + { + let nat_map = nat_map.clone(); + thread::spawn(move || { + if let Err(e) = tcp_proxy(tcp_listener, nat_map, stop_manager) { + log::warn!("tcp_proxy:{:?}", e); } - } + }); } + Ok(Self { port, nat_map }) } } -async fn proxy(client: TcpStream, server: TcpStream) -> io::Result<()> { - let (client_read, client_write) = client.into_split(); - let (server_read, server_write) = server.into_split(); - let time = Arc::new(AtomicCell::new(Instant::now())); - let time1 = time.clone(); - tokio::spawn(async move { - if let Err(e) = copy(client_read, server_write, &time1).await { - log::warn!("{:?}", e); - } - }); - copy(server_read, client_write, &time).await -} - -async fn copy( - mut read: OwnedReadHalf, - mut write: OwnedWriteHalf, - time: &AtomicCell, -) -> io::Result<()> { - let mut buf = [0; 10240]; - loop { - tokio::select! { - result = read.read(&mut buf) =>{ - let len = result?; - if len==0{ - break; - } - write.write_all(&buf[..len]).await?; - time.store(Instant::now()); - } - _ = tokio::time::sleep(Duration::from_secs(600)) =>{ - if time.load().elapsed()>=Duration::from_secs(580){ - //读写均超时再退出 - break; - } - } - } - } - Ok(()) -} - -#[derive(Clone)] -pub struct TcpHandler(u16, Arc>); - -impl ProxyHandler for TcpHandler { +impl ProxyHandler for TcpProxy { fn recv_handle( &self, ipv4: &mut IpV4Packet<&mut [u8]>, @@ -147,13 +59,14 @@ impl ProxyHandler for TcpHandler { let mut tcp_packet = TcpPacket::new(source, destination, ipv4.payload_mut())?; let source_port = tcp_packet.source_port(); let dest_port = tcp_packet.destination_port(); - tcp_packet.set_destination_port(self.0); + tcp_packet.set_destination_port(self.port); tcp_packet.update_checksum(); ipv4.set_destination_ip(destination); ipv4.update_checksum(); let key = SocketAddrV4::new(source, source_port); - //https://github.com/crossbeam-rs/crossbeam/issues/1023 - self.1.insert(key, SocketAddrV4::new(dest_ip, dest_port)); + self.nat_map + .lock() + .insert(key, SocketAddrV4::new(dest_ip, dest_port)); Ok(false) } @@ -164,8 +77,7 @@ impl ProxyHandler for TcpHandler { let tcp_packet = TcpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?; SocketAddrV4::new(dest_ip, tcp_packet.destination_port()) }; - if let Some(entry) = self.1.get(&dest_addr) { - let source_addr = entry.value(); + if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) { let source_ip = *source_addr.ip(); let mut tcp_packet = TcpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?; tcp_packet.set_source_port(source_addr.port()); @@ -176,3 +88,301 @@ impl ProxyHandler for TcpHandler { Ok(()) } } + +fn tcp_proxy( + mut tcp_listener: TcpListener, + nat_map: Arc>>, + stop_manager: StopManager, +) -> io::Result<()> { + let mut poll = Poll::new()?; + poll.registry() + .register(&mut tcp_listener, SERVER, Interest::READABLE)?; + let mut events = Events::with_capacity(32); + let mut tcp_map: HashMap = HashMap::with_capacity(16); + let mut mapping: HashMap = HashMap::with_capacity(16); + let stop = Waker::new(poll.registry(), NOTIFY)?; + let _worker = stop_manager.add_listener("tcp_proxy".into(), move || { + if let Err(e) = stop.wake() { + log::warn!("stop tcp_proxy:{:?}", e); + } + })?; + loop { + poll.poll(&mut events, None)?; + for event in events.iter() { + match event.token() { + SERVER => { + accept_handle( + poll.registry(), + &tcp_listener, + &nat_map, + &mut tcp_map, + &mut mapping, + ); + } + NOTIFY => { + return Ok(()); + } + Token(index) => { + let (val, src_index) = if let Some(v) = tcp_map.get_mut(&index) { + (v, index) + } else { + if let Some(dest_index) = mapping.get(&index) { + if let Some(v) = tcp_map.get_mut(dest_index) { + (v, *dest_index) + } else { + continue; + } + } else { + continue; + } + }; + let (stream1, stream2, buf1, buf2) = val.as_mut(index); + if event.is_readable() { + if let Err(e) = readable_handle(stream1, stream2, buf1) { + log::warn!("tcp proxy {:?}", e); + close(src_index, &mut tcp_map, &mut mapping); + } + } else if event.is_writable() { + let read = buf2.len() >= BUF_LEN; + if let Err(e) = writable_handle(stream1, buf2) { + log::warn!("tcp proxy {:?}", e); + close(src_index, &mut tcp_map, &mut mapping); + } else if read { + if let Err(e) = readable_handle(stream2, stream1, buf2) { + log::warn!("tcp proxy {:?}", e); + close(src_index, &mut tcp_map, &mut mapping); + } + } + } else { + close(src_index, &mut tcp_map, &mut mapping); + } + } + } + } + } +} + +fn accept_handle( + registry: &Registry, + tcp_listener: &TcpListener, + nat_map: &Mutex>, + tcp_map: &mut HashMap, + mapping: &mut HashMap, +) { + loop { + match tcp_listener.accept() { + Ok((mut src_stream, addr)) => { + #[cfg(windows)] + let src_fd = src_stream.as_raw_socket() as usize; + #[cfg(unix)] + let src_fd = src_stream.as_raw_fd() as usize; + if src_fd == SERVER_VAL || src_fd == NOTIFY_VAL { + log::error!("fd错误:{:?}", src_fd); + continue; + } + let addr = match addr { + SocketAddr::V4(addr) => addr, + SocketAddr::V6(_) => { + // 忽略ipv6 + continue; + } + }; + let _ = src_stream.set_nodelay(false); + if let Some(dest_addr) = nat_map.lock().get(&addr).cloned() { + match tcp_connect(addr.port(), dest_addr.into()) { + Ok(mut dest_stream) => { + #[cfg(windows)] + let dest_fd = dest_stream.as_raw_socket() as usize; + #[cfg(unix)] + let dest_fd = dest_stream.as_raw_fd() as usize; + if dest_fd == SERVER_VAL || dest_fd == NOTIFY_VAL { + log::error!("fd错误:{:?}", dest_fd); + continue; + } + if let Err(e) = registry.register( + &mut src_stream, + Token(src_fd), + Interest::READABLE, + ) { + log::error!("register src_stream:{:?}", e); + continue; + } + if let Err(e) = registry.register( + &mut dest_stream, + Token(dest_fd), + Interest::READABLE, + ) { + log::error!("register dest_stream:{:?}", e); + continue; + } + tcp_map.insert( + src_fd, + ProxyValue::new(src_stream, dest_stream, src_fd, dest_fd), + ); + mapping.insert(dest_fd, src_fd); + } + Err(e) => { + log::error!("connect:{:?}", e); + } + } + } + } + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + log::error!("accept:{:?}", e); + } + } + } +} + +fn tcp_connect(src_port: u16, addr: SocketAddr) -> io::Result { + let socket = socket2::Socket::new( + socket2::Domain::IPV4, + socket2::Type::STREAM, + Some(socket2::Protocol::TCP), + )?; + if socket + .bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into()) + .is_err() + { + socket.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?; + } + socket.set_nonblocking(true)?; + let _ = socket.set_nodelay(false); + if let Err(e) = socket.connect(&addr.into()) { + if e.kind() != io::ErrorKind::WouldBlock { + return Err(e); + } + } + Ok(TcpStream::from_std(socket.into())) +} + +struct ProxyValue { + src_stream: TcpStream, + dest_stream: TcpStream, + src_fd: usize, + dest_fd: usize, + src_buf: BytesMut, + dest_buf: BytesMut, +} + +const BUF_LEN: usize = 10 * 4096; + +impl ProxyValue { + fn new(src_stream: TcpStream, dest_stream: TcpStream, src_fd: usize, dest_fd: usize) -> Self { + Self { + src_stream, + dest_stream, + src_fd, + dest_fd, + src_buf: BytesMut::with_capacity(BUF_LEN), + dest_buf: BytesMut::with_capacity(BUF_LEN), + } + } + fn as_mut( + &mut self, + index: usize, + ) -> (&mut TcpStream, &mut TcpStream, &mut BytesMut, &mut BytesMut) { + if index == self.src_fd { + ( + &mut self.src_stream, + &mut self.dest_stream, + &mut self.src_buf, + &mut self.dest_buf, + ) + } else { + ( + &mut self.dest_stream, + &mut self.src_stream, + &mut self.dest_buf, + &mut self.src_buf, + ) + } + } +} + +fn readable_handle( + stream1: &mut TcpStream, + stream2: &mut TcpStream, + mid_buf: &mut BytesMut, +) -> io::Result<()> { + let mut buf = [0; BUF_LEN]; + + loop { + match stream1.read(&mut buf) { + Ok(len) => { + if len == 0 { + return Err(io::Error::from(io::ErrorKind::UnexpectedEof)); + } + let mut buf = &buf[..len]; + if mid_buf.is_empty() { + // 直接写入,避免在buf中过渡 + while !buf.is_empty() { + match stream2.write(buf) { + Ok(end) => { + if end == 0 { + return Err(io::Error::from(io::ErrorKind::WriteZero)); + } + buf = &buf[end..]; + } + Err(e) => { + if e.kind() != io::ErrorKind::WouldBlock { + return Err(e); + } + break; + } + } + } + if buf.is_empty() { + continue; + } + } + mid_buf.reserve(buf.len()); + mid_buf.put_slice(buf); + if mid_buf.len() >= BUF_LEN { + // 达到上限不再继续读取 + return Ok(()); + } + } + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + return Err(e); + } + } + } + Ok(()) +} + +fn writable_handle(stream: &mut TcpStream, mid_buf: &mut BytesMut) -> io::Result<()> { + while !mid_buf.is_empty() { + match stream.write(&mid_buf) { + Ok(len) => { + let _ = mid_buf.split_to(len); + } + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + return Err(e); + } + } + } + Ok(()) +} + +fn close( + index: usize, + tcp_map: &mut HashMap, + mapping: &mut HashMap, +) { + if let Some(val) = tcp_map.remove(&index) { + let _ = val.src_stream.shutdown(Shutdown::Both); + let _ = val.dest_stream.shutdown(Shutdown::Both); + mapping.remove(&val.src_fd); + mapping.remove(&val.dest_fd); + } +} diff --git a/vnt/src/ip_proxy/udp_proxy.rs b/vnt/src/ip_proxy/udp_proxy.rs index 89c3028..2ee2cd1 100644 --- a/vnt/src/ip_proxy/udp_proxy.rs +++ b/vnt/src/ip_proxy/udp_proxy.rs @@ -1,146 +1,56 @@ -use crate::ip_proxy::{DashMapNew, ProxyHandler}; -use crossbeam_utils::atomic::AtomicCell; -use dashmap::DashMap; +use std::net::{Ipv4Addr, SocketAddrV4}; +#[cfg(unix)] +use std::os::fd::AsRawFd; +#[cfg(windows)] +use std::os::windows::io::AsRawSocket; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use std::{collections::HashMap, io, net::SocketAddr, rc::Rc, thread}; + +use mio::{net::UdpSocket, Events, Interest, Poll, Token}; +use mio::{Registry, Waker}; +use parking_lot::Mutex; + use packet::ip::ipv4::packet::IpV4Packet; use packet::udp::udp::UdpPacket; -use std::io; -use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; -use std::sync::Arc; -use std::time::Duration; -use tokio::net::UdpSocket; -use tokio::time::Instant; -/// 一个udp代理,作用是利用系统协议栈,将udp数据报解析出来再转发到目的地址 +use crate::ip_proxy::ProxyHandler; +use crate::util::{Scheduler, StopManager}; + +const SERVER_VAL: usize = 0; +const SERVER: Token = Token(SERVER_VAL); +const NOTIFY_VAL: usize = 1; +const NOTIFY: Token = Token(NOTIFY_VAL); +// 开了ip代理后使用mstsc,mstsc会误以为在真实局域网,从而不维护udp心跳,导致断连,所以这里尽量长一点过期时间 +const NAT_TIMEOUT: Duration = Duration::from_secs(20 * 60); +const NAT_FAST_TIMEOUT: Duration = Duration::from_secs(5 * 60); +const NAT_MAX: usize = 5_000; + +#[derive(Clone)] pub struct UdpProxy { - udp_proxy_port: u16, - udp_socket: Arc, - map: Arc>, + port: u16, + nat_map: Arc>>, } impl UdpProxy { - pub fn new( - udp_socket: UdpSocket, - map: Arc>, - ) -> io::Result { - let udp_socket = Arc::new(udp_socket); - let udp_proxy_port = udp_socket.local_addr()?.port(); - Ok(Self { - udp_proxy_port, - udp_socket, - map, - }) - } - pub fn udp_handler(&self) -> UdpHandler { - UdpHandler(self.udp_proxy_port, self.map.clone()) - } - pub async fn start(self) { - let map = self.map; - let udp_socket = self.udp_socket; - let mut buf = [0u8; 65536]; - - let inner_map: Arc, Arc>)>> = - Arc::new(DashMap::new0()); - - loop { - match udp_socket.recv_from(&mut buf).await { - Ok((len, sender_addr)) => match sender_addr { - SocketAddr::V4(sender_addr) => { - match start0(&buf[..len], sender_addr, &inner_map, &map, &udp_socket).await - { - Ok(_) => {} - Err(e) => { - log::warn!("udp代理异常:{:?},来源:{}", e, sender_addr); - } - } - } - SocketAddr::V6(_) => {} - }, - Err(e) => { - log::warn!("udp代理异常:{:?}", e); - } - }; - } - } -} - -async fn start0( - buf: &[u8], - sender_addr: SocketAddrV4, - inner_map: &Arc, Arc>)>>, - map: &Arc>, - udp_socket: &Arc, -) -> io::Result<()> { - if let Some(entry) = inner_map.get(&sender_addr) { - entry.value().1.store(Instant::now()); - let udp = entry.value().0.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 = match UdpSocket::bind(format!("0.0.0.0:{}", sender_addr.port())).await + pub fn new(scheduler: Scheduler, stop_manager: StopManager) -> io::Result { + let nat_map: Arc>> = + Arc::new(Mutex::new(HashMap::with_capacity(16))); + let udp = UdpSocket::bind(format!("0.0.0.0:{}", 0).parse().unwrap())?; + let port = udp.local_addr()?.port(); { - Ok(udp) => udp, - Err(_) => UdpSocket::bind("0.0.0.0:0").await?, - }; - peer_udp_socket.connect(dest_addr).await?; - peer_udp_socket.send(buf).await?; - let peer_udp_socket = Arc::new(peer_udp_socket); - let inner_map = inner_map.clone(); - let time = Arc::new(AtomicCell::new(Instant::now())); - inner_map.insert(sender_addr, (peer_udp_socket.clone(), time.clone())); - let udp_socket = udp_socket.clone(); - let map = map.clone(); - tokio::spawn(async move { - let mut buf = [0u8; 65536]; - loop { - match tokio::time::timeout(Duration::from_secs(600), peer_udp_socket.recv(&mut buf)) - .await - { - Ok(rs) => match rs { - Ok(len) => match udp_socket.send_to(&buf[..len], sender_addr).await { - Ok(_) => {} - Err(e) => { - log::warn!( - "udp代理异常:{:?},来源:{},目标:{}", - e, - sender_addr, - dest_addr - ); - break; - } - }, - Err(e) => { - log::warn!( - "udp代理异常:{:?},来源:{},目标:{}", - e, - sender_addr, - dest_addr - ); - break; - } - }, - Err(_) => { - if time.load().elapsed() > Duration::from_secs(580) { - //超时关闭 - log::warn!("udp代理超时关闭,来源:{},目标:{}", sender_addr, dest_addr); - break; - } - } + let nat_map = nat_map.clone(); + thread::spawn(move || { + if let Err(e) = udp_proxy(udp, nat_map, scheduler, stop_manager) { + log::warn!("udp_proxy:{:?}", e); } - } - inner_map.remove(&sender_addr); - map.remove(&sender_addr); - }); + }); + } + Ok(Self { port, nat_map }) } - Ok(()) } -#[derive(Clone)] -pub struct UdpHandler(u16, Arc>); - -impl ProxyHandler for UdpHandler { +impl ProxyHandler for UdpProxy { fn recv_handle( &self, ipv4: &mut IpV4Packet<&mut [u8]>, @@ -152,12 +62,14 @@ impl ProxyHandler for UdpHandler { let mut udp_packet = UdpPacket::new(source, destination, ipv4.payload_mut())?; let source_port = udp_packet.source_port(); let dest_port = udp_packet.destination_port(); - udp_packet.set_destination_port(self.0); + udp_packet.set_destination_port(self.port); udp_packet.update_checksum(); ipv4.set_destination_ip(destination); ipv4.update_checksum(); let key = SocketAddrV4::new(source, source_port); - self.1.insert(key, SocketAddrV4::new(dest_ip, dest_port)); + self.nat_map + .lock() + .insert(key.into(), SocketAddrV4::new(dest_ip, dest_port).into()); Ok(false) } @@ -168,8 +80,7 @@ impl ProxyHandler for UdpHandler { let udp_packet = UdpPacket::new(src_ip, dest_ip, ipv4.payload_mut())?; SocketAddrV4::new(dest_ip, udp_packet.destination_port()) }; - if let Some(entry) = self.1.get(&dest_addr) { - let source_addr = entry.value(); + if let Some(source_addr) = self.nat_map.lock().get(&dest_addr) { let source_ip = *source_addr.ip(); let mut udp_packet = UdpPacket::new(source_ip, dest_ip, ipv4.payload_mut())?; udp_packet.set_source_port(source_addr.port()); @@ -180,3 +91,219 @@ impl ProxyHandler for UdpHandler { Ok(()) } } + +fn udp_proxy( + mut udp: UdpSocket, + nat_map: Arc>>, + scheduler: Scheduler, + stop_manager: StopManager, +) -> io::Result<()> { + let mut poll = Poll::new()?; + + poll.registry() + .register(&mut udp, SERVER, Interest::READABLE)?; + let mut events = Events::with_capacity(32); + let mut buf = [0; 65536]; + let mut token_map: HashMap, SocketAddrV4, Instant)> = + HashMap::with_capacity(64); + let mut udp_map: HashMap, Instant)> = HashMap::with_capacity(64); + let mut timeout = false; + let waker = Arc::new(Waker::new(poll.registry(), NOTIFY)?); + let stop = waker.clone(); + let _worker = stop_manager.add_listener("udp_proxy".into(), move || { + if let Err(e) = stop.wake() { + log::warn!("stop udp_proxy:{:?}", e); + } + })?; + loop { + let mut check = false; + if token_map.is_empty() { + poll.poll(&mut events, None)?; + } else { + //所有事件 50分钟超时 + if let Err(e) = poll.poll(&mut events, Some(Duration::from_secs(50 * 60))) { + if e.kind() == io::ErrorKind::TimedOut || e.kind() == io::ErrorKind::WouldBlock { + token_map.clear(); + udp_map.clear(); + continue; + } + return Err(e); + } + } + for event in events.iter() { + match event.token() { + SERVER => server_handle( + poll.registry(), + &udp, + &nat_map, + &mut token_map, + &mut udp_map, + &mut buf, + ), + NOTIFY => { + if stop_manager.is_stop() { + return Ok(()); + } + check = true; + } + token => { + if let Err(e) = readable_handle(&udp, &mut token_map, &token, &mut buf) { + log::error!("发送目标失败:{:?}", e); + if let Some((_, src_addr, _)) = token_map.remove(&token) { + udp_map.remove(&src_addr); + } + } + } + } + } + if check { + //超时校验 + if token_map.len() > NAT_MAX / 2 { + check_handle(&mut token_map, &mut udp_map, NAT_FAST_TIMEOUT) + } else { + check_handle(&mut token_map, &mut udp_map, NAT_TIMEOUT) + } + timeout = false; + } + if !token_map.is_empty() && !timeout { + //注册超时监听 + timeout = true; + let waker = waker.clone(); + scheduler.timeout(NAT_FAST_TIMEOUT, move |_| { + let _ = waker.wake(); + }); + } + } +} + +fn check_handle( + token_map: &mut HashMap, SocketAddrV4, Instant)>, + udp_map: &mut HashMap, Instant)>, + timeout: Duration, +) { + let mut remove_list = Vec::new(); + for (token, (_, addr, time)) in token_map.iter() { + if time.elapsed() > timeout { + if let Some((_, time)) = udp_map.get(addr) { + if time.elapsed() > timeout { + //映射超时,需要移除 + remove_list.push(*token); + } + } + } + } + for token in remove_list { + if let Some((_, src_addr, _)) = token_map.remove(&token) { + udp_map.remove(&src_addr); + } + } +} + +fn server_handle( + registry: &Registry, + udp: &UdpSocket, + nat_map: &Mutex>, + token_map: &mut HashMap, SocketAddrV4, Instant)>, + udp_map: &mut HashMap, Instant)>, + buf: &mut [u8], +) { + loop { + let (len, src_addr) = match udp.recv_from(buf) { + Ok((len, src_addr)) => match src_addr { + SocketAddr::V4(addr) => (len, addr), + SocketAddr::V6(_) => { + continue; + } + }, + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + log::error!("接收数据失败:{:?}", e); + break; + } + }; + if let Some((dest_udp, time)) = udp_map.get_mut(&src_addr) { + //发送失败就当丢包了 + let _ = dest_udp.send(&buf[..len]); + *time = Instant::now(); + } else if let Some(dest_addr) = nat_map.lock().get(&src_addr).cloned() { + if token_map.len() >= NAT_MAX { + log::error!( + "UDP NAT_MAX:src_addr={:?},dest_addr={:?}", + src_addr, + dest_addr + ); + continue; + } + match udp_connect(src_addr.port(), dest_addr.into()) { + Ok((token_val, mut dest_udp)) => { + let token = Token(token_val); + if let Err(e) = registry.register(&mut dest_udp, token, Interest::READABLE) { + log::error!("register失败:{:?},addr={:?}", e, dest_addr); + continue; + } + if dest_udp.send(&buf[..len]).is_ok() { + let dest_udp = Rc::new(dest_udp); + token_map.insert(token, (dest_udp.clone(), src_addr, Instant::now())); + udp_map.insert(src_addr, (dest_udp, Instant::now())); + } + } + Err(e) => { + log::error!("绑定目标地址失败:{:?}", e); + continue; + } + }; + } + } +} + +/// 得到一个 fd不为SERVER_VAL或者NOTYFY_VAL的socket +fn udp_connect(src_port: u16, addr: SocketAddr) -> io::Result<(usize, UdpSocket)> { + loop { + let udp = if let Ok(udp) = + UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into()) + { + udp + } else { + UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())? + }; + #[cfg(windows)] + let fd = udp.as_raw_socket() as usize; + #[cfg(unix)] + let fd = udp.as_raw_fd() as usize; + if fd == SERVER_VAL || fd == NOTIFY_VAL { + continue; + } + // 只接收目标的数据 + udp.connect(addr)?; + return Ok((fd, udp)); + } +} + +fn readable_handle( + udp: &UdpSocket, + token_map: &mut HashMap, SocketAddrV4, Instant)>, + token: &Token, + buf: &mut [u8], +) -> io::Result<()> { + if let Some((dest_udp, src_addr, time)) = token_map.get_mut(&token) { + loop { + let len = match dest_udp.recv(buf) { + Ok(rs) => rs, + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + return Err(e); + } + }; + if len == 0 { + return Err(io::Error::from(io::ErrorKind::UnexpectedEof)); + } + let _ = udp.send_to(&buf[..len], (*src_addr).into()); + } + *time = Instant::now(); + } + Ok(()) +}