From ff2bbdd8372f013249692191d471210fd8207ecc Mon Sep 17 00:00:00 2001 From: lbl8603 <49143209+lbl8603@users.noreply.github.com> Date: Wed, 8 May 2024 20:19:33 +0800 Subject: [PATCH] =?UTF-8?q?=E6=94=B9=E5=9B=9E=E7=94=A8tokio=E5=A4=84?= =?UTF-8?q?=E7=90=86=E4=BB=A3=E7=90=86=EF=BC=8C=E7=AE=80=E5=8C=96=E4=BB=A3?= =?UTF-8?q?=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 122 ++++++++++ vnt/Cargo.toml | 36 +-- vnt/src/ip_proxy/icmp_proxy.rs | 161 +++++-------- vnt/src/ip_proxy/mod.rs | 40 ++- vnt/src/ip_proxy/tcp_proxy.rs | 428 +++++---------------------------- vnt/src/ip_proxy/udp_proxy.rs | 337 ++++++++------------------ 6 files changed, 386 insertions(+), 738 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 18f0a8d..8a797cb 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,21 @@ # It is not intended for manual editing. version = 3 +[[package]] +name = "addr2line" +version = "0.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a30b2e23b9e17a9f90641c7ab1549cd9b44f296d3ccbf309d2863cfe398a0cb" +dependencies = [ + "gimli", +] + +[[package]] +name = "adler" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f26201604c87b1e01bd3d98f8d5d9a8fcbb815e8cedb41ffccbeb4bf593a35fe" + [[package]] name = "aead" version = "0.5.2" @@ -79,6 +94,21 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f1fdabc7756949593fe60f30ec81974b613357de856987752631dea1e3394c80" +[[package]] +name = "backtrace" +version = "0.3.71" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26b05800d2e817c8b3b4b54abd461726265fa9789ae34330622f2db9ee696f9d" +dependencies = [ + "addr2line", + "cc", + "cfg-if", + "libc", + "miniz_oxide", + "object", + "rustc-demangle", +] + [[package]] name = "base64ct" version = "1.6.0" @@ -422,6 +452,12 @@ dependencies = [ "polyval", ] +[[package]] +name = "gimli" +version = "0.28.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4271d37baee1b8c7e4b708028c57d816cf9d2434acb33a549475f78c181f6253" + [[package]] name = "hashbrown" version = "0.12.3" @@ -434,6 +470,12 @@ version = "0.14.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "290f1a1d9242c78d09ce40a5e87e7554ee637af1351968159f4952f028f75604" +[[package]] +name = "hermit-abi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d231dfb89cfffdbc30e7fc41579ed6066ad03abda9e567ccafae602b97ec5024" + [[package]] name = "home" version = "0.5.9" @@ -656,6 +698,15 @@ version = "2.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c8640c5d730cb13ebd907d8d04b52f55ac9a2eec55b440c8892f40d56c76c1d" +[[package]] +name = "miniz_oxide" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d811f3e15f28568be3407c8e7fdb6514c1cda3cb30683f15b6a1a1dc4ea14a7" +dependencies = [ + "adler", +] + [[package]] name = "mio" version = "0.8.11" @@ -726,6 +777,25 @@ dependencies = [ "libm", ] +[[package]] +name = "num_cpus" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4161fcb6d602d4d2081af7c3a45852d875a03dd337a6bfdd6e06407b61342a43" +dependencies = [ + "hermit-abi", + "libc", +] + +[[package]] +name = "object" +version = "0.32.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6a622008b6e321afc04970976f62ee297fdbaa6f95318ca343e3eebb9648441" +dependencies = [ + "memchr", +] + [[package]] name = "once_cell" version = "1.19.0" @@ -817,6 +887,12 @@ dependencies = [ "base64ct", ] +[[package]] +name = "pin-project-lite" +version = "0.2.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bda66fc9667c18cb2758a2ac84d1167245054bcf85d5d1aaa6923f45801bdd02" + [[package]] name = "pkcs1" version = "0.7.5" @@ -1096,6 +1172,12 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustc-demangle" +version = "0.1.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d626bb9dae77e28219937af045c257c28bfd3f69333c512553507f5f9798cb76" + [[package]] name = "rustix" version = "0.38.32" @@ -1195,6 +1277,15 @@ dependencies = [ "digest", ] +[[package]] +name = "signal-hook-registry" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9e9e0b4211b72e7b8b6e85c807d36c212bdb33ea8587f7569562a84df5465b1" +dependencies = [ + "libc", +] + [[package]] name = "signature" version = "2.2.0" @@ -1343,6 +1434,36 @@ dependencies = [ "winapi", ] +[[package]] +name = "tokio" +version = "1.37.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1adbebffeca75fcfd058afa480fb6c0b81e165a0323f9c9d39c9697e37c46787" +dependencies = [ + "backtrace", + "bytes", + "libc", + "mio", + "num_cpus", + "parking_lot", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys 0.48.0", +] + +[[package]] +name = "tokio-macros" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b8a1e28f2deaa14e508979454cb3a223b10b938b45af148bc0986de36f1923b" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.60", +] + [[package]] name = "tun" version = "0.1.0" @@ -1468,6 +1589,7 @@ dependencies = [ "spki", "stun-format", "thiserror", + "tokio", "tun", ] diff --git a/vnt/Cargo.toml b/vnt/Cargo.toml index a1e91c2..2ff4655 100644 --- a/vnt/Cargo.toml +++ b/vnt/Cargo.toml @@ -6,7 +6,7 @@ edition = "2021" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html [dependencies] -tun= {path = "tun"} +tun = { path = "tun" } packet = { path = "./packet" } bytes = "1.5.0" log = "0.4.17" @@ -19,22 +19,25 @@ sha2 = { version = "0.10.6", features = ["oid"] } thiserror = "1.0.37" protobuf = "3.2.0" socket2 = { version = "0.5.2", features = ["all"] } -aes-gcm = { version = "0.10.2",optional = true } +aes-gcm = { version = "0.10.2", optional = true } ring = { version = "0.17.0", optional = true } -cbc = {version = "0.1.2",optional = true} -ecb = {version = "0.1.2",optional = true} +cbc = { version = "0.1.2", optional = true } +ecb = { version = "0.1.2", optional = true } aes = "0.8.3" stun-format = { version = "1.0.1", features = ["fmt", "rfc3489"] } -rsa = { version = "0.9.2", features = [] ,optional = true} -spki = { version = "0.7.2", features = ["fingerprint", "alloc","base64"] ,optional = true} -openssl-sys = { git = "https://github.com/lbl8603/rust-openssl" ,optional = true} -libsm = {git="https://github.com/lbl8603/libsm" ,optional = true} +rsa = { version = "0.9.2", features = [], optional = true } +spki = { version = "0.7.2", features = ["fingerprint", "alloc", "base64"], optional = true } +openssl-sys = { git = "https://github.com/lbl8603/rust-openssl", optional = true } +libsm = { git = "https://github.com/lbl8603/libsm", optional = true } -mio = {version = "0.8.10",features = ["os-poll","net"]} +mio = { version = "0.8.10", features = ["os-poll", "net"] } crossbeam-queue = "0.3.11" anyhow = "1.0.82" dns-parser = "0.8.0" +tokio = { version = "1.37.0", features = ["full"], optional = true } + + [target.'cfg(target_os = "windows")'.dependencies] libloading = "0.8.0" @@ -44,14 +47,15 @@ protobuf-codegen = "3.2.0" protoc-bin-vendored = "3.0.0" [features] -default = ["server_encrypt","aes_gcm","aes_cbc","aes_ecb","sm4_cbc","ip_proxy"] +default = ["server_encrypt", "aes_gcm", "aes_cbc", "aes_ecb", "sm4_cbc", "ip_proxy", "port_mapping"] openssl = ["openssl-sys"] # 从源码编译 openssl-vendored = ["openssl-sys/vendored"] ring-cipher = ["ring"] -aes_cbc=["cbc"] -aes_ecb=["ecb"] -sm4_cbc=["libsm"] -aes_gcm=["aes-gcm"] -server_encrypt =["aes-gcm","rsa","spki"] -ip_proxy=[] +aes_cbc = ["cbc"] +aes_ecb = ["ecb"] +sm4_cbc = ["libsm"] +aes_gcm = ["aes-gcm"] +server_encrypt = ["aes-gcm", "rsa", "spki"] +ip_proxy = ["tokio"] +port_mapping = ["tokio"] diff --git a/vnt/src/ip_proxy/icmp_proxy.rs b/vnt/src/ip_proxy/icmp_proxy.rs index d24dbd4..7c4f68f 100644 --- a/vnt/src/ip_proxy/icmp_proxy.rs +++ b/vnt/src/ip_proxy/icmp_proxy.rs @@ -1,12 +1,12 @@ +use anyhow::Context; use std::collections::HashMap; +use std::io; use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}; use std::sync::Arc; -use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; -use mio::net::UdpSocket; -use mio::{Events, Interest, Poll, Token, Waker}; use parking_lot::Mutex; +use tokio::net::UdpSocket; use packet::icmp::icmp; use packet::icmp::icmp::HeaderOther; @@ -18,7 +18,6 @@ use crate::handle::CurrentDeviceInfo; use crate::ip_proxy::ProxyHandler; use crate::protocol; use crate::protocol::{NetPacket, MAX_TTL}; -use crate::util::StopManager; #[derive(Clone)] pub struct IcmpProxy { icmp_socket: Arc, @@ -27,48 +26,50 @@ pub struct IcmpProxy { } impl IcmpProxy { - pub fn new( + pub async fn new( context: ChannelContext, - stop_manager: StopManager, current_device: Arc>, client_cipher: Cipher, - ) -> io::Result { + ) -> anyhow::Result { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] let icmp_socket = socket2::Socket::new( socket2::Domain::IPV4, socket2::Type::RAW, Some(socket2::Protocol::ICMPV4), - )?; + ) + .context("new Socket RAW ICMPV4 failed")?; #[cfg(target_os = "android")] let icmp_socket = socket2::Socket::new( socket2::Domain::IPV4, socket2::Type::DGRAM, Some(socket2::Protocol::ICMPV4), - )?; + ) + .context("new Socket DGRAM ICMPV4 failed")?; let addr: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); - icmp_socket.bind(&socket2::SockAddr::from(addr))?; + icmp_socket + .bind(&socket2::SockAddr::from(addr)) + .context("bind Socket ICMPV4 failed")?; 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 tokio_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::Builder::new() - .name("icmpProxy".into()) - .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); - } - }) - .expect("icmpProxy"); + tokio::spawn(async { + if let Err(e) = icmp_proxy( + tokio_icmp_socket, + nat_map, + context, + current_device, + client_cipher, + ) + .await + { + log::warn!("icmp_proxy:{:?}", e); + } + }); } Ok(Self { icmp_socket: Arc::new(std_socket), @@ -77,105 +78,51 @@ impl IcmpProxy { } } -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, +async fn icmp_proxy( + icmp_socket: UdpSocket, // 对端-> 真实来源 nat_map: Arc>>, context: ChannelContext, - 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 = Arc::new(Waker::new(poll.registry(), NOTIFY)?); - let _stop = stop.clone(); - 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)?; - if stop_manager.is_stop() { - return Ok(()); - } - for event in events.iter() { - match event.token() { - SERVER => readable_handle( - &icmp_socket, + #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] + let start = 12; + #[cfg(target_os = "android")] + let start = 12 + 20; + loop { + let (len, addr) = icmp_socket.recv_from(&mut buf[start..]).await?; + if let IpAddr::V4(peer_ip) = addr.ip() { + #[cfg(target_os = "android")] + { + let buf = &mut buf[12..]; + // ipv4 头部20字节 + buf[0] = 0b0100_0110; + //写入总长度 + buf[2..4].copy_from_slice(&((20 + len) as u16).to_be_bytes()); + + let mut ipv4 = IpV4Packet::unchecked(buf); + ipv4.set_flags(2); + ipv4.set_ttl(1); + ipv4.set_protocol(packet::ip::ipv4::protocol::Protocol::Icmp); + ipv4.set_source_ip(peer_ip); + } + recv_handle( &mut buf, + start + len, + peer_ip, &nat_map, &context, ¤t_device, &client_cipher, - ), - NOTIFY => { - return Ok(()); - } - _ => {} + ); } } } } -fn readable_handle( - icmp_socket: &UdpSocket, - buf: &mut [u8], - nat_map: &Mutex>, - context: &ChannelContext, - current_device: &AtomicCell, - client_cipher: &Cipher, -) { - #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] - let start = 12; - #[cfg(target_os = "android")] - let start = 12 + 20; - loop { - let (len, addr) = match icmp_socket.recv_from(&mut buf[start..]) { - Ok(rs) => rs, - Err(e) => { - if e.kind() == io::ErrorKind::WouldBlock { - break; - } - log::warn!("icmp_socket {:?}", e); - return; - } - }; - if let IpAddr::V4(peer_ip) = addr.ip() { - #[cfg(target_os = "android")] - { - let buf = &mut buf[12..]; - // ipv4 头部20字节 - buf[0] = 0b0100_0110; - //写入总长度 - buf[2..4].copy_from_slice(&((20 + len) as u16).to_be_bytes()); - let mut ipv4 = IpV4Packet::unchecked(buf); - ipv4.set_flags(2); - ipv4.set_ttl(1); - ipv4.set_protocol(packet::ip::ipv4::protocol::Protocol::Icmp); - ipv4.set_source_ip(peer_ip); - } - recv_handle( - buf, - start + len, - peer_ip, - &nat_map, - &context, - ¤t_device, - &client_cipher, - ); - } - } -} fn recv_handle( buf: &mut [u8], data_len: usize, diff --git a/vnt/src/ip_proxy/mod.rs b/vnt/src/ip_proxy/mod.rs index 54a97c5..2387aef 100644 --- a/vnt/src/ip_proxy/mod.rs +++ b/vnt/src/ip_proxy/mod.rs @@ -1,6 +1,6 @@ -use std::io; use std::net::Ipv4Addr; use std::sync::Arc; +use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; @@ -13,7 +13,7 @@ use crate::handle::CurrentDeviceInfo; 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}; +use crate::util::StopManager; pub mod icmp_proxy; pub mod tcp_proxy; @@ -38,14 +38,40 @@ pub struct IpProxyMap { pub fn init_proxy( context: ChannelContext, - scheduler: Scheduler, stop_manager: StopManager, current_device: Arc>, client_cipher: Cipher, -) -> 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)?; +) -> anyhow::Result { + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .thread_name("ipProxy") + .build()?; + let proxy_map = runtime.block_on(init_proxy0(context, current_device, client_cipher))?; + let (sender, receiver) = tokio::sync::oneshot::channel::<()>(); + let worker = stop_manager.add_listener("ipProxy".into(), move || { + let _ = sender.send(()); + })?; + thread::Builder::new() + .name("ipProxy".into()) + .spawn(move || { + runtime.block_on(async { + let _ = receiver.await; + }); + runtime.shutdown_background(); + drop(worker); + })?; + + return Ok(proxy_map); +} + +async fn init_proxy0( + context: ChannelContext, + current_device: Arc>, + client_cipher: Cipher, +) -> anyhow::Result { + let icmp_proxy = IcmpProxy::new(context, current_device, client_cipher).await?; + let tcp_proxy = TcpProxy::new().await?; + let udp_proxy = UdpProxy::new().await?; Ok(IpProxyMap { icmp_proxy, diff --git a/vnt/src/ip_proxy/tcp_proxy.rs b/vnt/src/ip_proxy/tcp_proxy.rs index 5c460e2..0e0c123 100644 --- a/vnt/src/ip_proxy/tcp_proxy.rs +++ b/vnt/src/ip_proxy/tcp_proxy.rs @@ -1,28 +1,16 @@ -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 anyhow::Context; +use std::net::{Ipv4Addr, SocketAddrV4}; use std::sync::Arc; use std::time::Duration; -use std::{collections::HashMap, io, net::SocketAddr, thread}; +use std::{collections::HashMap, io, net::SocketAddr}; -use bytes::{BufMut, BytesMut}; -use mio::net::TcpStream; -use mio::{net::TcpListener, Events, Interest, Poll, Registry, Token, Waker}; use parking_lot::Mutex; +use tokio::net::{TcpListener, TcpSocket, TcpStream}; use packet::ip::ipv4::packet::IpV4Packet; use packet::tcp::tcp::TcpPacket; 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 { @@ -31,21 +19,16 @@ pub struct TcpProxy { } impl TcpProxy { - pub fn new(stop_manager: StopManager) -> io::Result { + pub async fn new() -> anyhow::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 tcp_listener = TcpListener::bind(format!("0.0.0.0:{}", 0)) + .await + .context("TcpProxy bind failed")?; let port = tcp_listener.local_addr()?.port(); { let nat_map = nat_map.clone(); - thread::Builder::new() - .name("tcpProxy".into()) - .spawn(move || { - if let Err(e) = tcp_proxy(tcp_listener, nat_map, stop_manager) { - log::warn!("tcp_proxy:{:?}", e); - } - }) - .expect("tcpProxy"); + tokio::spawn(tcp_proxy(tcp_listener, nat_map)); } Ok(Self { port, nat_map }) } @@ -93,365 +76,74 @@ impl ProxyHandler for TcpProxy { } } -fn tcp_proxy( - mut tcp_listener: TcpListener, +async fn tcp_proxy( + 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 = Arc::new(Waker::new(poll.registry(), NOTIFY)?); - let _stop = stop.clone(); - 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)?; - if stop_manager.is_stop() { - return Ok(()); - } - 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, state1, state2) = val.as_mut(index); - if event.is_readable() { - if let Err(_) = readable_handle(stream1, stream2, buf1, state2) { - *state1 |= READ_CLOSED; - } - } - if event.is_writable() { - let read = buf2.len() >= BUF_LEN; - if let Err(_) = writable_handle(stream1, buf2) { - *state1 |= WRITE_CLOSED; - } else if read { - if readable_handle(stream2, stream1, buf2, state1).is_err() { - *state2 |= READ_CLOSED; - } - } - } - if event.is_read_closed() || event.is_error() { - *state1 |= READ_CLOSED; - } - if event.is_write_closed() || event.is_error() { - *state1 |= WRITE_CLOSED; - } - if is_write_closed(*state1) { - let _ = stream1.shutdown(Shutdown::Write); - let _ = stream2.shutdown(Shutdown::Read); - } - if is_read_closed(*state1) { - let _ = stream1.shutdown(Shutdown::Read); - if buf1.is_empty() { - let _ = stream2.shutdown(Shutdown::Write); - } - } - if (is_both_closed(*state1) && buf1.is_empty()) - || (is_both_closed(*state2) && buf2.is_empty()) - || (is_write_closed(*state1) && is_write_closed(*state2) - || (is_read_closed(*state1) - && is_read_closed(*state2) - && buf1.is_empty() - && buf2.is_empty())) - { - 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.add(Interest::WRITABLE), - ) { - log::error!("register src_stream:{:?}", e); - continue; - } - if let Err(e) = registry.register( - &mut dest_stream, - Token(dest_fd), - Interest::READABLE.add(Interest::WRITABLE), - ) { - 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, addr, dest_addr); - } + match tcp_listener.accept().await { + Ok((tcp_stream, sender_addr)) => match sender_addr { + SocketAddr::V4(sender_addr) => { + if let Some(dest_addr) = nat_map.lock().get(&sender_addr).cloned() { + tokio::spawn(async move { + let peer_tcp_stream = + match tcp_connect(sender_addr.port(), dest_addr.into()).await { + Ok(peer_tcp_stream) => peer_tcp_stream, + Err(e) => { + log::warn!( + "tcp代理异常:{:?},来源:{},目标:{}", + e, + sender_addr, + dest_addr + ); + return; + } + }; + proxy(sender_addr, dest_addr, tcp_stream, peer_tcp_stream).await + }); + } else { + log::warn!("tcp代理异常: 来源:{},未找到目标", sender_addr); } } - } + SocketAddr::V6(_) => {} + }, Err(e) => { - if e.kind() == io::ErrorKind::WouldBlock { - break; - } - log::error!("accept:{:?}", e); + log::warn!("tcp代理监听:{:?}", 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), - )?; +/// 优先使用来源端口建立tcp连接 +async fn tcp_connect(src_port: u16, addr: SocketAddr) -> anyhow::Result { + let socket = TcpSocket::new_v4()?; if socket - .bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into()) + .bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into()) .is_err() { - socket.bind(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?; - } - if let Err(e) = socket.set_tcp_keepalive( - &socket2::TcpKeepalive::new() - .with_time(Duration::from_secs(120)) - .with_interval(Duration::from_secs(10)), - ) { - log::warn!("set_tcp_keepalive err {:?}", e); + socket.bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?; } let _ = socket.set_nodelay(false); - socket.connect_timeout(&addr.into(), Duration::from_secs(3))?; - socket.set_nonblocking(true)?; - Ok(TcpStream::from_std(socket.into())) + let tcp_stream = tokio::time::timeout(Duration::from_secs(5), socket.connect(addr)) + .await + .with_context(|| format!("TCP connection timeout {}", addr))? + .with_context(|| format!("TCP connection target failed {}", addr))?; + Ok(tcp_stream) } -#[derive(Debug)] -struct ProxyValue { - src_stream: TcpStream, - dest_stream: TcpStream, - src_fd: usize, - dest_fd: usize, - src_buf: BytesMut, - dest_buf: BytesMut, - src_state: u8, - dest_state: u8, -} - -const BUF_LEN: usize = 65536; - -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), - src_state: NORMAL, - dest_state: NORMAL, - } - } - fn as_mut( - &mut self, - index: usize, - ) -> ( - &mut TcpStream, - &mut TcpStream, - &mut BytesMut, - &mut BytesMut, - &mut u8, - &mut u8, - ) { - if index == self.src_fd { - ( - &mut self.src_stream, - &mut self.dest_stream, - &mut self.src_buf, - &mut self.dest_buf, - &mut self.src_state, - &mut self.dest_state, - ) - } else { - ( - &mut self.dest_stream, - &mut self.src_stream, - &mut self.dest_buf, - &mut self.src_buf, - &mut self.dest_state, - &mut self.src_state, - ) - } - } -} - -fn readable_handle( - stream1: &mut TcpStream, - stream2: &mut TcpStream, - mid_buf: &mut BytesMut, - state2: &mut u8, -) -> io::Result<()> { - let mut buf = [0; BUF_LEN]; - - loop { - if mid_buf.len() >= BUF_LEN { - // 达到上限不再继续读取 - return Ok(()); - } - 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 { - *state2 |= WRITE_CLOSED; - return Err(io::Error::from(io::ErrorKind::WriteZero)); - } - buf = &buf[end..]; - } - Err(e) => { - if e.kind() != io::ErrorKind::WouldBlock { - *state2 |= WRITE_CLOSED; - return Err(e); - } - break; - } - } - } - if buf.is_empty() { - continue; - } - } - mid_buf.reserve(buf.len()); - mid_buf.put_slice(buf); - } - 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, +async fn proxy( + sender_addr: SocketAddrV4, + dest_addr: SocketAddrV4, + client: TcpStream, + server: TcpStream, ) { - 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); + let (mut client_read, mut client_write) = client.into_split(); + let (mut server_read, mut server_write) = server.into_split(); + tokio::spawn(async move { + if let Err(e) = tokio::io::copy(&mut client_read, &mut server_write).await { + log::warn!("client tcp proxy {}->{},{:?}", sender_addr, dest_addr, e); + } + }); + if let Err(e) = tokio::io::copy(&mut server_read, &mut client_write).await { + log::warn!("server tcp proxy {}->{},{:?}", sender_addr, dest_addr, e); } } - -const NORMAL: u8 = 0b00; -const READ_CLOSED: u8 = 0b01; -const WRITE_CLOSED: u8 = 0b10; -const BOTH_CLOSED: u8 = 0b11; - -fn is_read_closed(state: u8) -> bool { - (state & READ_CLOSED == READ_CLOSED) || is_both_closed(state) -} - -fn is_write_closed(state: u8) -> bool { - (state & WRITE_CLOSED == WRITE_CLOSED) || is_both_closed(state) -} - -fn is_both_closed(state: u8) -> bool { - state & BOTH_CLOSED == BOTH_CLOSED -} diff --git a/vnt/src/ip_proxy/udp_proxy.rs b/vnt/src/ip_proxy/udp_proxy.rs index e416aee..0837749 100644 --- a/vnt/src/ip_proxy/udp_proxy.rs +++ b/vnt/src/ip_proxy/udp_proxy.rs @@ -1,30 +1,17 @@ +use anyhow::Context; +use crossbeam_utils::atomic::AtomicCell; 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 std::{collections::HashMap, io, net::SocketAddr}; -use mio::{net::UdpSocket, Events, Interest, Poll, Token}; -use mio::{Registry, Waker}; use parking_lot::Mutex; +use tokio::net::UdpSocket; use packet::ip::ipv4::packet::IpV4Packet; use packet::udp::udp::UdpPacket; 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 { @@ -33,21 +20,20 @@ pub struct UdpProxy { } impl UdpProxy { - pub fn new(scheduler: Scheduler, stop_manager: StopManager) -> io::Result { + pub async fn new() -> anyhow::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 udp = UdpSocket::bind(format!("0.0.0.0:{}", 0)) + .await + .context("UdpProxy bind failed")?; let port = udp.local_addr()?.port(); { let nat_map = nat_map.clone(); - thread::Builder::new() - .name("udpProxy".into()) - .spawn(move || { - if let Err(e) = udp_proxy(udp, nat_map, scheduler, stop_manager) { - log::warn!("udp_proxy:{:?}", e); - } - }) - .expect("udpProxy"); + tokio::spawn(async { + if let Err(e) = udp_proxy(udp, nat_map).await { + log::warn!("udp_proxy:{:?}", e); + } + }); } Ok(Self { port, nat_map }) } @@ -95,230 +81,101 @@ impl ProxyHandler for UdpProxy { } } -fn udp_proxy( - mut udp: UdpSocket, +async fn udp_proxy( + udp: UdpSocket, nat_map: Arc>>, - scheduler: Scheduler, - stop_manager: StopManager, ) -> io::Result<()> { - let mut poll = Poll::new()?; + let mut buf = [0u8; 65536]; - 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); - } - })?; + let inner_map: Arc, Arc>)>>> = + Arc::new(Mutex::new(HashMap::with_capacity(64))); + let udp_socket = Arc::new(udp); 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 { - log::warn!( - "超时清理所有udp映射 {},token_map={},udp_map={}", - e, - token_map.len(), - udp_map.len() - ); - token_map.clear(); - udp_map.clear(); - continue; + match udp_socket.recv_from(&mut buf).await { + Ok((len, sender_addr)) => match sender_addr { + SocketAddr::V4(sender_addr) => { + if let Err(e) = + udp_proxy0(&buf[..len], sender_addr, &inner_map, &nat_map, &udp_socket) + .await + { + log::warn!("udp proxy {} {:?}", sender_addr, e); + } } - return Err(e); + SocketAddr::V6(_) => {} + }, + Err(e) => { + log::warn!("udp代理异常:{:?}", e); } - } - if stop_manager.is_stop() { - return Ok(()); - } - 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 => { - 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); + }; + } +} + +async fn udp_proxy0( + buf: &[u8], + sender_addr: SocketAddrV4, + inner_map: &Arc, Arc>)>>>, + map: &Arc>>, + udp_socket: &Arc, +) -> io::Result<()> { + let option = inner_map.lock().get(&sender_addr).cloned(); + if let Some((udp, time)) = option { + time.store(Instant::now()); + udp.send(buf).await?; + } else { + let option = map.lock().get(&sender_addr).cloned(); + if let Some(dest_addr) = option { + //先使用相同的端口,冲突了再随机端口 + let peer_udp_socket = + match UdpSocket::bind(format!("0.0.0.0:{}", sender_addr.port())).await { + 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 + .lock() + .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 proxy {}->{} {:?}", sender_addr, dest_addr, e); + break; + } + }, + Err(e) => { + log::warn!("udp proxy {}->{} {:?}", sender_addr, dest_addr, e); + + break; + } + }, + Err(_) => { + if time.load().elapsed() > Duration::from_secs(580) { + //超时关闭 + log::warn!("udp proxy timeout {}->{}", sender_addr, dest_addr); + break; + } } } } - } - } - 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(); + inner_map.lock().remove(&sender_addr); + map.lock().remove(&sender_addr); }); } } -} - -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); - log::warn!( - "超时清理udp映射 {},token_map={},udp_map={}", - src_addr, - token_map.len(), - udp_map.len() - ); - } - } -} - -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; - } - let _ = dest_udp.send(&buf[..len]); - 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(()) }