diff --git a/Cargo.lock b/Cargo.lock index 10045db..5f57b7c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -858,9 +858,9 @@ checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" [[package]] name = "libc" -version = "0.2.153" +version = "0.2.155" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9c198f91728a82281a64e1f4f9eeb25d82cb32a5de251c6bd1b5154d63a8e7bd" +checksum = "97b3888a4aecf77e811145cadf6eef5901f4782c53886191b2f693f24761847c" [[package]] name = "libloading" @@ -1005,6 +1005,18 @@ dependencies = [ "windows-sys 0.48.0", ] +[[package]] +name = "network-interface" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "433419f898328beca4f2c6c73a1b52540658d92b0a99f0269330457e0fd998d5" +dependencies = [ + "cc", + "libc", + "thiserror", + "winapi", +] + [[package]] name = "nom" version = "7.1.3" @@ -1747,9 +1759,9 @@ checksum = "3c5e1a9a646d36c3599cd173a41282daf47c44583ad367b8e6837255952e5c67" [[package]] name = "socket2" -version = "0.5.6" +version = "0.5.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "05ffd9c0a93b7543e062e759284fcf5f5e3b098501104bfbdde4d404db792871" +checksum = "ce305eb0b4296696835b71df73eb912e0f1ffd2556a501fcede6e0c50349191c" dependencies = [ "libc", "windows-sys 0.52.0", @@ -2153,6 +2165,7 @@ dependencies = [ "log", "lz4_flex", "mio", + "network-interface", "openssl-sys", "packet", "parking_lot", @@ -2171,6 +2184,7 @@ dependencies = [ "tokio", "tokio-tungstenite", "tun", + "windows-sys 0.52.0", "zstd", ] diff --git a/common/src/cli.rs b/common/src/cli.rs index d340ade..4cf6aee 100644 --- a/common/src/cli.rs +++ b/common/src/cli.rs @@ -76,6 +76,7 @@ pub fn parse_args_config() -> anyhow::Result, bool)> opts.optmulti("", "vnt-mapping", "vnt-mapping", ""); opts.optopt("f", "", "配置文件", ""); opts.optopt("", "compressor", "压缩算法", ""); + opts.optopt("", "local-ipv4", "指定本地ipv4网卡IP", ""); opts.optflag("", "disable-stats", "关闭流量统计"); opts.optflag("", "allow-wg", "允许接入WireGuard"); //"后台运行时,查看其他设备列表" @@ -283,6 +284,15 @@ pub fn parse_args_config() -> anyhow::Result, bool)> #[cfg(feature = "port_mapping")] let port_mapping_list = matches.opt_strs("mapping"); let vnt_mapping_list = matches.opt_strs("vnt-mapping"); + let local_ipv4: Option = matches.opt_get("local-ipv4").unwrap(); + let local_ipv4 = local_ipv4 + .map(|v| Ipv4Addr::from_str(&v).expect(&format!("'--local-ipv4 {}' error", v))); + if let Some(local_ipv4) = local_ipv4 { + if local_ipv4.is_unspecified() || local_ipv4.is_broadcast() || local_ipv4.is_multicast() + { + return Err(anyhow::anyhow!("'--local-ipv4 {}' invalid", local_ipv4)); + } + } let disable_stats = matches.opt_present("disable-stats"); let allow_wire_guard = matches.opt_present("allow-wg"); let compressor = if let Some(compressor) = matches.opt_str("compressor").as_ref() { @@ -326,6 +336,7 @@ pub fn parse_args_config() -> anyhow::Result, bool)> compressor, !disable_stats, allow_wire_guard, + local_ipv4, ) { Ok(config) => config, Err(e) => { @@ -378,6 +389,7 @@ fn get_description(key: &str, language: &str) -> String { ("--compressor-lz4 ", ("启用压缩,可选值lz4,例如 --compressor lz4", "Enable compression, option lz4, e.g., --compressor lz4")), ("--compressor-zstd ", ("启用压缩,可选值zstd<,level>,level为压缩级别,例如 --compressor zstd,10", "Enable compression, options zstd<,level>, level is compression level, e.g., --compressor zstd,10")), ("--vnt-mapping ", ("vnt地址映射,例如 --vnt-mapping tcp:80-10.26.0.10:80 映射目标是vnt网络或其子网中的设备", "VNT address mapping, e.g., --vnt-mapping tcp:80-10.26.0.10:80 maps to a device in VNT network or its subnet")), + ("--local-ipv4", ("本地出口网卡的ipv4地址", "IPv4 address of local export network card")), ("--disable-stats", ("关闭流量统计", "Disable traffic statistics")), ("--list", ("后台运行时,查看其他设备列表", "View list of other devices when running in background")), ("--all", ("后台运行时,查看其他设备完整信息", "View complete information of other devices when running in background")), @@ -564,6 +576,10 @@ fn print_usage(program: &str, _opts: Options) { " --vnt-mapping {}", green(get_description("--vnt-mapping ", &language).to_string()) ); + println!( + " --local-ipv4 {}", + get_description("--local-ipv4", &language) + ); println!( " --disable-stats {}", get_description("--disable-stats", &language) diff --git a/common/src/config/file_config.rs b/common/src/config/file_config.rs index c5b2b1b..0a8c0b7 100644 --- a/common/src/config/file_config.rs +++ b/common/src/config/file_config.rs @@ -48,6 +48,7 @@ pub struct FileConfig { pub disable_stats: bool, // 允许传递wg流量 pub allow_wire_guard: bool, + pub local_ipv4: Option, } impl Default for FileConfig { @@ -93,6 +94,7 @@ impl Default for FileConfig { vnt_mapping: vec![], disable_stats: false, allow_wire_guard: false, + local_ipv4: None, } } } @@ -181,6 +183,7 @@ pub fn read_config(file_path: &str) -> anyhow::Result<(Config, Vec, bool compressor, !file_conf.disable_stats, file_conf.allow_wire_guard, + file_conf.local_ipv4, )?; Ok((config, file_conf.vnt_mapping, file_conf.cmd)) diff --git a/vnt/Cargo.toml b/vnt/Cargo.toml index 63d8c8a..a7cfa8f 100644 --- a/vnt/Cargo.toml +++ b/vnt/Cargo.toml @@ -18,7 +18,7 @@ rand = "0.8.5" sha2 = { version = "0.10.6", features = ["oid"] } thiserror = "1.0.37" protobuf = "=3.2.0" -socket2 = { version = "0.5.2", features = ["all"] } +socket2 = { version = "0.5.7", features = ["all"] } aes-gcm = { version = "0.10.2", optional = true } ring = { version = "0.17.0", optional = true } cbc = { version = "0.1.2", optional = true } @@ -46,10 +46,17 @@ fnv = "1.0.7" igd = { version = "0.12.1", optional = true } tokio-tungstenite = { version = "0.23.1", optional = true } rustls = { version = "0.23.0", features = ["ring"], default-features = false, optional = true } + +network-interface = "2.0.0" + futures-util = "0.3.30" [target.'cfg(target_os = "windows")'.dependencies] libloading = "0.8.0" - +windows-sys = {version = "0.52.0",features = [ "Win32_Foundation", + "Win32_Networking_WinSock", + "Win32_System_IO", + "Win32_System_Threading", + "Win32_System_WindowsProgramming",]} [build-dependencies] protobuf-codegen = "=3.2.0" diff --git a/vnt/src/channel/context.rs b/vnt/src/channel/context.rs index 0bdce08..505df6f 100644 --- a/vnt/src/channel/context.rs +++ b/vnt/src/channel/context.rs @@ -1,7 +1,7 @@ use fnv::FnvHashMap; -use std::net::{Ipv4Addr, SocketAddr, SocketAddrV6, UdpSocket}; +use std::net::{Ipv4Addr, SocketAddr, UdpSocket}; use std::ops::Deref; -use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::atomic::AtomicUsize; use std::sync::Arc; use std::time::{Duration, Instant}; use std::{io, thread}; @@ -12,6 +12,7 @@ use rand::Rng; use crate::channel::punch::NatType; use crate::channel::sender::{AcceptSocketSender, PacketSender}; +use crate::channel::socket::LocalInterface; use crate::channel::{ConnectProtocol, Route, RouteKey, UseChannelType, DEFAULT_RT}; use crate::protocol::NetPacket; use crate::util::limit::TrafficMeterMultiAddress; @@ -25,16 +26,17 @@ pub struct ChannelContext { impl ChannelContext { pub fn new( main_udp_socket: Vec, + v4_len: usize, use_channel_type: UseChannelType, first_latency: bool, protocol: ConnectProtocol, packet_loss_rate: Option, packet_delay: u32, - use_ipv6: bool, up_traffic_meter: Option, down_traffic_meter: Option, + default_interface: LocalInterface, ) -> Self { - let channel_num = main_udp_socket.len(); + let channel_num = v4_len; assert_ne!(channel_num, 0, "not channel"); let packet_loss_rate = packet_loss_rate .map(|v| { @@ -48,16 +50,16 @@ impl ChannelContext { .unwrap_or(0); let inner = ContextInner { main_udp_socket, + v4_len, sub_udp_socket: RwLock::new(Vec::new()), packet_map: RwLock::new(FnvHashMap::default()), route_table: RouteTable::new(use_channel_type, first_latency, channel_num), protocol, packet_loss_rate, packet_delay, - main_index: AtomicUsize::new(0), - use_ipv6, up_traffic_meter, down_traffic_meter, + default_interface, }; Self { inner: Arc::new(inner), @@ -80,6 +82,7 @@ const PACKET_LOSS_RATE_DENOMINATOR: u32 = 100_0000; pub struct ContextInner { // 核心udp socket pub(crate) main_udp_socket: Vec, + v4_len: usize, // 对称网络增加的udp socket sub_udp_socket: RwLock>, // tcp数据发送器 @@ -92,16 +95,18 @@ pub struct ContextInner { packet_loss_rate: u32, //控制延迟 packet_delay: u32, - main_index: AtomicUsize, - use_ipv6: bool, pub(crate) up_traffic_meter: Option, pub(crate) down_traffic_meter: Option, + default_interface: LocalInterface, } impl ContextInner { pub fn use_channel_type(&self) -> UseChannelType { self.route_table.use_channel_type } + pub fn default_interface(&self) -> &LocalInterface { + &self.default_interface + } /// 通过sub_udp_socket是否为空来判断是否为锥形网络 pub fn is_cone(&self) -> bool { self.sub_udp_socket.read().is_empty() @@ -120,7 +125,7 @@ impl ContextInner { &self, nat_type: NatType, udp_socket_sender: &AcceptSocketSender>>, - ) -> io::Result<()> { + ) -> anyhow::Result<()> { let mut write_guard = self.sub_udp_socket.write(); match nat_type { NatType::Symmetric => { @@ -129,9 +134,11 @@ impl ContextInner { } let mut vec = Vec::with_capacity(SYMMETRIC_CHANNEL_NUM); for _ in 0..SYMMETRIC_CHANNEL_NUM { - let udp = UdpSocket::bind("0.0.0.0:0")?; - //副通道使用异步io - udp.set_nonblocking(true)?; + let udp = crate::channel::socket::bind_udp( + "0.0.0.0:0".parse().unwrap(), + &self.default_interface, + )?; + let udp: UdpSocket = udp.into(); vec.push(udp); } let mut mio_vec = Vec::with_capacity(SYMMETRIC_CHANNEL_NUM); @@ -152,14 +159,18 @@ impl ContextInner { } Ok(()) } - + #[inline] pub fn channel_num(&self) -> usize { + self.v4_len + } + #[inline] + pub fn main_len(&self) -> usize { self.main_udp_socket.len() } /// 获取核心udp监听的端口,用于其他客户端连接 pub fn main_local_udp_port(&self) -> io::Result> { let mut ports = Vec::new(); - for udp in self.main_udp_socket.iter() { + for udp in self.main_udp_socket[..self.v4_len].iter() { ports.push(udp.local_addr()?.port()) } Ok(ports) @@ -171,20 +182,13 @@ impl ContextInner { Err(io::Error::from(io::ErrorKind::NotFound)) } } - pub fn send_main_udp(&self, index: usize, buf: &[u8], mut addr: SocketAddr) -> io::Result<()> { - if self.use_ipv6 { - //如果是v4地址则需要转换成v6 - if let SocketAddr::V4(ipv4) = addr { - addr = SocketAddr::V6(SocketAddrV6::new( - ipv4.ip().to_ipv6_mapped(), - ipv4.port(), - 0, - 0, - )); - } + pub fn send_main_udp(&self, index: usize, buf: &[u8], addr: SocketAddr) -> io::Result<()> { + if let Some(udp) = self.main_udp_socket.get(index) { + udp.send_to(buf, addr)?; + Ok(()) + } else { + Err(io::Error::new(io::ErrorKind::Other, "overflow")) } - self.main_udp_socket[index].send_to(buf, addr)?; - Ok(()) } /// 将数据发送到默认通道,一般发往服务器才用此方法 pub fn send_default>( @@ -193,7 +197,11 @@ impl ContextInner { addr: SocketAddr, ) -> io::Result<()> { if self.protocol.is_udp() { - self.send_main_udp(self.main_index.load(Ordering::Relaxed), buf.buffer(), addr)? + if addr.is_ipv4() { + self.send_main_udp(0, buf.buffer(), addr)? + } else { + self.send_main_udp(self.v4_len, buf.buffer(), addr)? + } } else { self.send_tcp(buf.buffer(), addr)? } @@ -203,10 +211,6 @@ impl ContextInner { Ok(()) } - pub fn change_main_index(&self) { - let index = (self.main_index.load(Ordering::Relaxed) + 1) % self.main_udp_socket.len(); - self.main_index.store(index, Ordering::Relaxed); - } /// 此方法仅用于对称网络打洞 pub fn try_send_all(&self, buf: &[u8], addr: SocketAddr) { self.try_send_all_main(buf, addr); @@ -287,7 +291,7 @@ impl ContextInner { if let Some(udp) = self .sub_udp_socket .read() - .get(route_key.index - self.main_udp_socket.len()) + .get(route_key.index - self.main_len()) { udp.send_to(buf.buffer(), route_key.addr)?; } else { diff --git a/vnt/src/channel/mod.rs b/vnt/src/channel/mod.rs index 639ab4c..8cc1f3b 100644 --- a/vnt/src/channel/mod.rs +++ b/vnt/src/channel/mod.rs @@ -6,6 +6,7 @@ use tokio::sync::mpsc::channel; use crate::channel::context::ChannelContext; use crate::channel::handler::RecvChannelHandler; use crate::channel::sender::{AcceptSocketSender, ConnectUtil}; +use crate::channel::socket::{bind_udp, LocalInterface}; use crate::channel::tcp_channel::tcp_listen; use crate::channel::udp_channel::udp_listen; #[cfg(feature = "ws")] @@ -19,6 +20,7 @@ pub mod idle; pub mod notify; pub mod punch; pub mod sender; +pub mod socket; pub mod tcp_channel; pub mod udp_channel; #[cfg(feature = "ws")] @@ -201,11 +203,13 @@ pub(crate) fn init_context( protocol: ConnectProtocol, packet_loss_rate: Option, packet_delay: u32, + default_interface: LocalInterface, up_traffic_meter: Option, down_traffic_meter: Option, ) -> anyhow::Result<(ChannelContext, std::net::TcpListener)> { assert!(!ports.is_empty(), "not channel"); - let mut udps = Vec::with_capacity(ports.len()); + let mut main_udp_socket_v4 = Vec::with_capacity(ports.len()); + let mut main_udp_socket_v6 = Vec::with_capacity(ports.len()); //检查系统是否支持ipv6 let use_ipv6 = match socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None) { Ok(_) => true, @@ -215,40 +219,33 @@ pub(crate) fn init_context( } }; for port in &ports { - //监听v6+v4双栈 - let (socket, address) = if use_ipv6 { - let address: SocketAddr = format!("[::]:{}", port).parse().unwrap(); - let socket = socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None)?; - socket - .set_only_v6(false) - .with_context(|| format!("set_only_v6 failed: {}", &address))?; - (socket, address) + let addr_v4: SocketAddr = format!("0.0.0.0:{}", port).parse().unwrap(); + if use_ipv6 { + let (main_channel_v4, main_channel_v6) = bind_udp_v4_and_v6(*port, &default_interface)?; + main_udp_socket_v4.push(main_channel_v4); + main_udp_socket_v6.push(main_channel_v6); } else { - let address: SocketAddr = format!("0.0.0.0:{}", port).parse().unwrap(); - ( - socket2::Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None)?, - address, - ) - }; - if let Err(e) = socket.set_recv_buffer_size(2 * 1024 * 1024) { - log::warn!("set_recv_buffer_size {:?}", e); + let socket = bind_udp(addr_v4, &default_interface)?; + let main_channel_v4: UdpSocket = socket.into(); + main_udp_socket_v4.push(main_channel_v4); } - socket - .bind(&address.into()) - .with_context(|| format!("bind failed: {}", &address))?; - let main_channel: UdpSocket = socket.into(); - udps.push(main_channel); } + let mut main_udp_socket = + Vec::with_capacity(main_udp_socket_v4.len() + main_udp_socket_v6.len()); + let v4_len = main_udp_socket_v4.len(); + main_udp_socket.append(&mut main_udp_socket_v4); + main_udp_socket.append(&mut main_udp_socket_v6); let context = ChannelContext::new( - udps, + main_udp_socket, + v4_len, use_channel_type, first_latency, protocol, packet_loss_rate, packet_delay, - use_ipv6, up_traffic_meter, down_traffic_meter, + default_interface, ); let port = context.main_local_udp_port()?[0]; @@ -265,7 +262,7 @@ pub(crate) fn init_context( let socket = socket2::Socket::new(socket2::Domain::IPV4, socket2::Type::STREAM, None)?; (socket, address) }; - + let _ = socket.set_reuse_address(true); if let Err(e) = socket.bind(&address.into()) { if ports[0] == 0 { //端口可能冲突,则使用任意端口 @@ -285,9 +282,49 @@ pub(crate) fn init_context( } socket.listen(128)?; socket.set_nonblocking(true)?; - socket.set_nodelay(false)?; + socket.set_nodelay(true)?; Ok((context, socket.into())) } +fn bind_udp_v4_and_v6( + port: u16, + default_interface: &LocalInterface, +) -> anyhow::Result<(UdpSocket, UdpSocket)> { + let mut count = 0; + loop { + let addr_v4: SocketAddr = format!("0.0.0.0:{}", port).parse().unwrap(); + let socket = bind_udp(addr_v4, default_interface)?; + if let Err(e) = socket.set_recv_buffer_size(2 * 1024 * 1024) { + log::warn!("set_recv_buffer_size {:?}", e); + } + let main_channel_v4: UdpSocket = socket.into(); + let addr = main_channel_v4.local_addr()?; + let addr_v6: SocketAddr = format!("[::]:{}", addr.port()).parse().unwrap(); + let socket = if port == 0 { + match bind_udp(addr_v6, default_interface) { + Ok(socket) => socket, + Err(e) => { + if count > 10 { + return Err(e); + } + if let Some(e) = e.downcast_ref::() { + if e.kind() == std::io::ErrorKind::AddrInUse { + count += 1; + continue; + } + } + Err(e)? + } + } + } else { + bind_udp(addr_v6, default_interface)? + }; + if let Err(e) = socket.set_recv_buffer_size(2 * 1024 * 1024) { + log::warn!("set_recv_buffer_size {:?}", e); + } + let main_channel_v6: UdpSocket = socket.into(); + return Ok((main_channel_v4, main_channel_v6)); + } +} pub(crate) fn init_channel( tcp_listener: std::net::TcpListener, diff --git a/vnt/src/channel/punch.rs b/vnt/src/channel/punch.rs index 07764a8..fdc22b9 100644 --- a/vnt/src/channel/punch.rs +++ b/vnt/src/channel/punch.rs @@ -1,4 +1,3 @@ -use crossbeam_utils::atomic::AtomicCell; use std::collections::HashMap; use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; use std::ops::{Div, Mul}; @@ -7,12 +6,12 @@ use std::sync::Arc; use std::time::Duration; use std::{io, thread}; +use crossbeam_utils::atomic::AtomicCell; use rand::prelude::SliceRandom; use rand::Rng; use crate::channel::context::ChannelContext; use crate::channel::sender::ConnectUtil; -use crate::external_route::ExternalRoute; use crate::handle::CurrentDeviceInfo; use crate::nat::NatTest; @@ -189,7 +188,6 @@ pub struct Punch { punch_model: PunchModel, is_tcp: bool, connect_util: ConnectUtil, - external_route: ExternalRoute, nat_test: NatTest, current_device: Arc>, } @@ -200,7 +198,6 @@ impl Punch { punch_model: PunchModel, is_tcp: bool, connect_util: ConnectUtil, - external_route: ExternalRoute, nat_test: NatTest, current_device: Arc>, ) -> Self { @@ -215,7 +212,6 @@ impl Punch { punch_model, is_tcp, connect_util, - external_route, nat_test, current_device, } @@ -242,19 +238,12 @@ impl Punch { return Ok(()); } let device_info = self.current_device.load(); - nat_info.public_ips.retain(|ip| { - self.external_route.route(ip).is_none() && device_info.not_in_network(*ip) - }); - nat_info.local_ipv4.filter(|ip| { - self.external_route.route(ip).is_none() && device_info.not_in_network(*ip) - }); - nat_info.ipv6.filter(|ip| { - if let Some(ip) = ip.to_ipv4() { - self.external_route.route(&ip).is_none() - } else { - true - } - }); + nat_info + .public_ips + .retain(|ip| device_info.not_in_network(*ip)); + nat_info + .local_ipv4 + .filter(|ip| device_info.not_in_network(*ip)); if punch_tcp && self.is_tcp && nat_info.tcp_port != 0 { //向tcp发起连接 if let Some(ipv6_addr) = nat_info.local_tcp_ipv6addr() { @@ -270,6 +259,7 @@ impl Punch { } } let channel_num = self.context.channel_num(); + let main_len = self.context.main_len(); for index in 0..channel_num { if let Some(ipv4_addr) = nat_info.local_udp_ipv4addr(index) { if !self.nat_test.is_local_address(false, ipv4_addr) { @@ -279,7 +269,7 @@ impl Punch { } if self.punch_model != PunchModel::IPv4 { - for index in 0..channel_num { + for index in channel_num..main_len { if let Some(ipv6_addr) = nat_info.local_udp_ipv6addr(index) { if !self.nat_test.is_local_address(false, ipv6_addr) { let rs = self.context.send_main_udp(index, buf, ipv6_addr); diff --git a/vnt/src/channel/socket/mod.rs b/vnt/src/channel/socket/mod.rs new file mode 100644 index 0000000..14052c2 --- /dev/null +++ b/vnt/src/channel/socket/mod.rs @@ -0,0 +1,111 @@ +use anyhow::{anyhow, Context}; +use network_interface::{NetworkInterface, NetworkInterfaceConfig}; +use socket2::Protocol; +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +#[cfg(unix)] +pub use unix::*; +#[cfg(windows)] +pub use windows::*; + +#[cfg(unix)] +mod unix; +#[cfg(windows)] +mod windows; + +pub trait VntSocketTrait { + fn set_ip_unicast_if(&self, _interface: &LocalInterface) -> anyhow::Result<()> { + Ok(()) + } +} + +#[derive(Clone, Debug, Default)] +pub struct LocalInterface { + index: u32, + #[cfg(unix)] + name: Option, +} + +pub async fn connect_tcp( + addr: SocketAddr, + default_interface: &LocalInterface, +) -> anyhow::Result { + let socket = create_tcp(addr.is_ipv4(), default_interface)?; + Ok(socket.connect(addr).await?) +} + +pub fn create_tcp( + v4: bool, + default_interface: &LocalInterface, +) -> anyhow::Result { + let socket = if v4 { + socket2::Socket::new( + socket2::Domain::IPV4, + socket2::Type::STREAM, + Some(Protocol::TCP), + )? + } else { + socket2::Socket::new( + socket2::Domain::IPV6, + socket2::Type::STREAM, + Some(Protocol::TCP), + )? + }; + if v4 { + socket.set_ip_unicast_if(default_interface)?; + } + socket.set_nonblocking(true)?; + socket.set_nodelay(true)?; + Ok(tokio::net::TcpSocket::from_std_stream(socket.into())) +} +pub fn bind_udp_ops( + addr: SocketAddr, + only_v6: bool, + default_interface: &LocalInterface, +) -> anyhow::Result { + let socket = if addr.is_ipv4() { + let socket = socket2::Socket::new( + socket2::Domain::IPV4, + socket2::Type::DGRAM, + Some(Protocol::UDP), + )?; + socket.set_ip_unicast_if(default_interface)?; + socket + } else { + let socket = socket2::Socket::new( + socket2::Domain::IPV6, + socket2::Type::DGRAM, + Some(Protocol::UDP), + )?; + socket + .set_only_v6(only_v6) + .with_context(|| format!("set_only_v6 failed: {}", &addr))?; + socket + }; + socket.set_nonblocking(true)?; + socket.bind(&addr.into())?; + Ok(socket) +} +pub fn bind_udp( + addr: SocketAddr, + default_interface: &LocalInterface, +) -> anyhow::Result { + bind_udp_ops(addr, true, default_interface).with_context(|| format!("{}", addr)) +} + +pub fn get_interface(dest_ip: Ipv4Addr) -> anyhow::Result { + let network_interfaces = NetworkInterface::show()?; + for iface in network_interfaces { + for addr in iface.addr { + if let IpAddr::V4(ip) = addr.ip() { + if ip == dest_ip { + return Ok(LocalInterface { + index: iface.index, + #[cfg(unix)] + name: Some(iface.name), + }); + } + } + } + } + Err(anyhow!("No network card with IP {} found", dest_ip)) +} diff --git a/vnt/src/channel/socket/unix.rs b/vnt/src/channel/socket/unix.rs new file mode 100644 index 0000000..d638cf6 --- /dev/null +++ b/vnt/src/channel/socket/unix.rs @@ -0,0 +1,35 @@ +use crate::channel::socket::{get_interface, LocalInterface, VntSocketTrait}; +use anyhow::Context; +use std::net::Ipv4Addr; + +#[cfg(any(target_os = "linux", target_os = "android"))] +impl VntSocketTrait for socket2::Socket { + fn set_ip_unicast_if(&self, interface: &LocalInterface) -> anyhow::Result<()> { + if let Some(name) = &interface.name { + self.bind_device(Some(name.as_bytes())) + .context("bind_device")?; + } + Ok(()) + } +} +#[cfg(target_os = "macos")] +impl VntSocketTrait for socket2::Socket { + fn set_ip_unicast_if(&self, interface: &LocalInterface) -> anyhow::Result<()> { + if interface.index != 0 { + self.bind_device_by_index_v4(std::num::NonZeroU32::new(interface.index)) + .with_context(|| format!("bind_device_by_index_v4 {:?}", interface))?; + } + Ok(()) + } +} + +pub fn get_best_interface(dest_ip: Ipv4Addr) -> anyhow::Result { + match get_interface(dest_ip) { + Ok(iface) => return Ok(iface), + Err(e) => { + log::warn!("not find interface e={:?},ip={}", e, dest_ip); + } + } + // 应该再查路由表找到默认路由的 + Ok(LocalInterface::default()) +} diff --git a/vnt/src/channel/socket/windows.rs b/vnt/src/channel/socket/windows.rs new file mode 100644 index 0000000..a518178 --- /dev/null +++ b/vnt/src/channel/socket/windows.rs @@ -0,0 +1,58 @@ +use std::mem; +use std::net::Ipv4Addr; +use std::os::windows::io::AsRawSocket; + +use windows_sys::core::PCSTR; +use windows_sys::Win32::NetworkManagement::IpHelper::GetBestInterfaceEx; +use windows_sys::Win32::Networking::WinSock::{ + htonl, setsockopt, AF_INET, IPPROTO_IP, IP_UNICAST_IF, SOCKADDR, SOCKADDR_IN, SOCKET_ERROR, +}; + +use crate::channel::socket::{LocalInterface, VntSocketTrait}; + +impl VntSocketTrait for socket2::Socket { + fn set_ip_unicast_if(&self, interface: &LocalInterface) -> anyhow::Result<()> { + let index = interface.index; + if index == 0 { + return Ok(()); + } + let raw_socket = self.as_raw_socket(); + let result = unsafe { + let best_interface = htonl(index); + setsockopt( + raw_socket as usize, + IPPROTO_IP, + IP_UNICAST_IF, + &best_interface as *const _ as PCSTR, + mem::size_of_val(&best_interface) as i32, + ) + }; + if result == SOCKET_ERROR { + Err(anyhow::anyhow!( + "Failed to set IP_UNICAST_IF: {:?} {}", + std::io::Error::last_os_error(), + index + ))?; + } + Ok(()) + } +} + +pub fn get_best_interface(dest_ip: Ipv4Addr) -> anyhow::Result { + // 获取最佳接口 + let index = unsafe { + let mut dest: SOCKADDR_IN = mem::zeroed(); + dest.sin_family = AF_INET as u16; + dest.sin_addr.S_un.S_addr = u32::from_ne_bytes(dest_ip.octets()); + + let mut index: u32 = 0; + if GetBestInterfaceEx(&dest as *const _ as *mut SOCKADDR, &mut index) != 0 { + Err(anyhow::anyhow!( + "Failed to GetBestInterfaceEx: {:?}", + std::io::Error::last_os_error() + ))?; + } + index + }; + Ok(LocalInterface { index }) +} diff --git a/vnt/src/channel/tcp_channel.rs b/vnt/src/channel/tcp_channel.rs index 97debc9..71aadd6 100644 --- a/vnt/src/channel/tcp_channel.rs +++ b/vnt/src/channel/tcp_channel.rs @@ -87,8 +87,11 @@ async fn connect_tcp0( where H: RecvChannelHandler, { - let mut stream = - tokio::time::timeout(Duration::from_secs(3), TcpStream::connect(addr)).await??; + let mut stream = tokio::time::timeout( + Duration::from_secs(3), + crate::channel::socket::connect_tcp(addr, context.default_interface()), + ) + .await??; tcp_write(&mut stream, &data).await?; tcp_stream_handle(stream, addr, recv_handler, context).await; diff --git a/vnt/src/channel/udp_channel.rs b/vnt/src/channel/udp_channel.rs index a895e6f..213de42 100644 --- a/vnt/src/channel/udp_channel.rs +++ b/vnt/src/channel/udp_channel.rs @@ -1,4 +1,3 @@ -use std::collections::HashMap; use std::sync::mpsc::{sync_channel, Receiver}; use std::{io, thread}; @@ -71,7 +70,8 @@ where let mut events = Events::with_capacity(1024); let mut buf = [0; BUFFER_SIZE]; let mut extend = [0; BUFFER_SIZE]; - let mut read_map: HashMap = HashMap::with_capacity(32); + let mut list: Vec = Vec::with_capacity(100); + let main_len = context.main_len(); loop { if let Err(e) = poll.poll(&mut events, None) { crate::ignore_io_interrupted(e)?; @@ -88,39 +88,43 @@ where match option { None => { log::info!("切换成锥形模式"); - for (_, mut udp_socket) in read_map.drain() { + for mut udp_socket in list.drain(..) { if let Err(e) = udp_socket.deregister(poll.registry()) { log::error!("{:?}", e); } } } Some(socket_list) => { + for mut udp_socket in list.drain(..) { + if let Err(e) = udp_socket.deregister(poll.registry()) { + log::error!("deregister {:?}", e); + } + } log::info!("切换成对称模式 监听端口数:{}", socket_list.len()); for (index, mut udp_socket) in socket_list.into_iter().enumerate() { - let token = Token(index + context.channel_num()); poll.registry().register( &mut udp_socket, - token, + Token(index), Interest::READABLE, )?; - read_map.insert(token, udp_socket); + list.push(udp_socket); } } } } } } - token => { - if let Some(udp_socket) = read_map.get(&token) { + Token(index) => { + if let Some(udp_socket) = list.get(index) { loop { match udp_socket.recv_from(&mut buf) { Ok((len, addr)) => { recv_handler.handle( &mut buf[..len], &mut extend, - RouteKey::new(ConnectProtocol::UDP, token.0, addr), + RouteKey::new(ConnectProtocol::UDP, index + main_len, addr), &context, ); } @@ -266,6 +270,7 @@ where for x in events.iter() { let index = match x.token() { NOTIFY => return Ok(()), + // 0的位置留给NOTIFY了,这里要再减回去,因为路由是通过index来找到对应udp的 Token(index) => index - 1, }; let udp = if let Some(udp) = udps.get(index) { diff --git a/vnt/src/core/conn.rs b/vnt/src/core/conn.rs index cacd907..261f9eb 100644 --- a/vnt/src/core/conn.rs +++ b/vnt/src/core/conn.rs @@ -12,6 +12,7 @@ use crate::channel::context::ChannelContext; use crate::channel::idle::Idle; use crate::channel::punch::{NatInfo, Punch}; use crate::channel::sender::IpPacketSender; +use crate::channel::socket::LocalInterface; use crate::channel::{init_channel, init_context, Route, RouteKey}; use crate::cipher::Cipher; #[cfg(feature = "server_encrypt")] @@ -29,7 +30,7 @@ use crate::tun_tap_device::tun_create_helper::{DeviceAdapter, TunDeviceHelper}; use crate::tun_tap_device::vnt_device::DeviceWrite; use crate::util::limit::TrafficMeterMultiAddress; use crate::util::{Scheduler, StopManager}; -use crate::{nat, VntCallback}; +use crate::{channel, nat, VntCallback}; #[derive(Clone)] pub struct Vnt { @@ -105,6 +106,7 @@ impl VntInner { } else { (None, None) }; + //服务端非对称加密 #[cfg(feature = "server_encrypt")] let rsa_cipher: Arc>> = Arc::new(Mutex::new(None)); @@ -131,6 +133,23 @@ impl VntInner { //设备列表 let device_map: Arc)>> = Arc::new(Mutex::new((0, HashMap::with_capacity(16)))); + let local_ipv4 = if let Some(local_ipv4) = config.local_ipv4 { + Some(local_ipv4) + } else { + nat::local_ipv4() + }; + + let default_interface = if config.in_ips.is_empty() { + //没有改变路由,不需要绑定网卡 + LocalInterface::default() + } else { + //vnt的流量都走这个接口 + let default_interface = + channel::socket::get_best_interface(local_ipv4.unwrap_or(Ipv4Addr::UNSPECIFIED))?; + log::info!("default_interface = {:?}", default_interface); + default_interface + }; + //基础信息 let config_info = BaseConfigInfo::new( config.name.clone(), @@ -149,6 +168,7 @@ impl VntInner { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] config.device_name.clone(), config.allow_wire_guard, + default_interface.clone(), ); // 服务停止管理器 let stop_manager = { @@ -179,10 +199,10 @@ impl VntInner { config.protocol, config.packet_loss_rate, config.packet_delay, + default_interface, up_traffic_meter.clone(), down_traffic_meter.clone(), )?; - let local_ipv4 = nat::local_ipv4(); let local_ipv6 = nat::local_ipv6(); let udp_ports = context.main_local_udp_port()?; let tcp_port = tcp_listener.local_addr()?.port(); @@ -194,6 +214,7 @@ impl VntInner { local_ipv6, udp_ports, tcp_port, + config.local_ipv4.is_none(), ); // 定时器 let scheduler = Scheduler::new(stop_manager.clone())?; @@ -268,7 +289,6 @@ impl VntInner { config.punch_model, config.protocol.is_base_tcp(), connect_util.clone(), - external_route.clone(), nat_test.clone(), current_device.clone(), ); diff --git a/vnt/src/core/mod.rs b/vnt/src/core/mod.rs index 56701ad..ac2d54b 100644 --- a/vnt/src/core/mod.rs +++ b/vnt/src/core/mod.rs @@ -52,6 +52,7 @@ pub struct Config { pub compressor: Compressor, pub enable_traffic: bool, pub allow_wire_guard: bool, + pub local_ipv4: Option, } impl Config { @@ -91,6 +92,7 @@ impl Config { enable_traffic: bool, // 允许传递wg流量 allow_wire_guard: bool, + local_ipv4: Option, ) -> anyhow::Result { for x in stun_server.iter_mut() { if !x.contains(":") { @@ -147,6 +149,9 @@ impl Config { *dest = *mask & *dest; } in_ips.sort_by(|(dest1, _, _), (dest2, _, _)| dest2.cmp(dest1)); + if let Some(local_ip) = local_ipv4 { + let _ = crate::channel::socket::get_interface(local_ip)?; + } Ok(Self { #[cfg(feature = "integrated_tun")] #[cfg(target_os = "windows")] @@ -184,6 +189,7 @@ impl Config { compressor, enable_traffic, allow_wire_guard, + local_ipv4, }) } } diff --git a/vnt/src/handle/maintain/re_nat_type.rs b/vnt/src/handle/maintain/re_nat_type.rs index 0c784fc..ca6d29f 100644 --- a/vnt/src/handle/maintain/re_nat_type.rs +++ b/vnt/src/handle/maintain/re_nat_type.rs @@ -29,9 +29,13 @@ fn retrieve_nat_type0( .name("natTest".into()) .spawn(move || { if nat_test.can_update() { - let local_ipv4 = nat::local_ipv4(); + let local_ipv4 = if nat_test.update_local_ipv4 { + nat::local_ipv4() + } else { + None + }; let local_ipv6 = nat::local_ipv6(); - match nat_test.re_test(local_ipv4, local_ipv6) { + match nat_test.re_test(local_ipv4, local_ipv6, context.default_interface()) { Ok(nat_info) => { log::info!("当前nat信息:{:?}", nat_info); if let Err(e) = context.switch(nat_info.nat_type, &udp_socket_sender) { diff --git a/vnt/src/handle/mod.rs b/vnt/src/handle/mod.rs index 754b7b3..ea51732 100644 --- a/vnt/src/handle/mod.rs +++ b/vnt/src/handle/mod.rs @@ -1,3 +1,4 @@ +use crate::channel::socket::LocalInterface; use crossbeam_utils::atomic::AtomicCell; use std::net::{IpAddr, Ipv4Addr, SocketAddr}; @@ -70,6 +71,7 @@ pub struct BaseConfigInfo { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] pub device_name: Option, pub allow_wire_guard: bool, + pub default_interface: LocalInterface, } impl BaseConfigInfo { @@ -90,6 +92,7 @@ impl BaseConfigInfo { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] device_name: Option, allow_wire_guard: bool, + default_interface: LocalInterface, ) -> Self { Self { name, @@ -108,6 +111,7 @@ impl BaseConfigInfo { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] device_name, allow_wire_guard, + default_interface, } } } diff --git a/vnt/src/handle/recv_data/server.rs b/vnt/src/handle/recv_data/server.rs index 2abce5c..f1b9407 100644 --- a/vnt/src/handle/recv_data/server.rs +++ b/vnt/src/handle/recv_data/server.rs @@ -105,12 +105,11 @@ impl PacketHandler for ServerPacketHandl ) -> anyhow::Result<()> { if !current_device.is_server_addr(route_key.addr) { //拦截不是服务端的流量 - log::info!( + log::warn!( "route_key={:?},不是来源于服务端地址{}", route_key, current_device.connect_server ); - return Ok(()); } context .route_table diff --git a/vnt/src/ip_proxy/icmp_proxy.rs b/vnt/src/ip_proxy/icmp_proxy.rs index ffb7d2e..60353df 100644 --- a/vnt/src/ip_proxy/icmp_proxy.rs +++ b/vnt/src/ip_proxy/icmp_proxy.rs @@ -13,6 +13,7 @@ use packet::icmp::icmp::HeaderOther; use packet::ip::ipv4::packet::IpV4Packet; use crate::channel::context::ChannelContext; +use crate::channel::socket::{LocalInterface, VntSocketTrait}; use crate::cipher::Cipher; use crate::handle::CurrentDeviceInfo; use crate::ip_proxy::ProxyHandler; @@ -30,6 +31,7 @@ impl IcmpProxy { context: ChannelContext, current_device: Arc>, client_cipher: Cipher, + default_interface: &LocalInterface, ) -> anyhow::Result { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] let icmp_socket = socket2::Socket::new( @@ -50,6 +52,7 @@ impl IcmpProxy { .bind(&socket2::SockAddr::from(addr)) .context("bind Socket ICMPV4 failed")?; icmp_socket.set_nonblocking(true)?; + icmp_socket.set_ip_unicast_if(default_interface)?; let std_socket: std::net::UdpSocket = icmp_socket.into(); let tokio_icmp_socket = UdpSocket::from_std(std_socket.try_clone()?)?; diff --git a/vnt/src/ip_proxy/mod.rs b/vnt/src/ip_proxy/mod.rs index 0361889..6f16916 100644 --- a/vnt/src/ip_proxy/mod.rs +++ b/vnt/src/ip_proxy/mod.rs @@ -68,14 +68,16 @@ pub fn init_proxy( } async fn init_proxy0( - _context: ChannelContext, + context: ChannelContext, _current_device: Arc>, _client_cipher: Cipher, ) -> anyhow::Result { + let default_interface = context.default_interface().clone(); #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] - let icmp_proxy = IcmpProxy::new(_context, _current_device, _client_cipher).await?; - let tcp_proxy = TcpProxy::new().await?; - let udp_proxy = UdpProxy::new().await?; + let icmp_proxy = + IcmpProxy::new(context, _current_device, _client_cipher, &default_interface).await?; + let tcp_proxy = TcpProxy::new(default_interface.clone()).await?; + let udp_proxy = UdpProxy::new(default_interface.clone()).await?; Ok(IpProxyMap { #[cfg(any(target_os = "windows", target_os = "linux", target_os = "macos"))] diff --git a/vnt/src/ip_proxy/tcp_proxy.rs b/vnt/src/ip_proxy/tcp_proxy.rs index 0e0c123..333f8df 100644 --- a/vnt/src/ip_proxy/tcp_proxy.rs +++ b/vnt/src/ip_proxy/tcp_proxy.rs @@ -5,13 +5,13 @@ use std::time::Duration; use std::{collections::HashMap, io, net::SocketAddr}; use parking_lot::Mutex; -use tokio::net::{TcpListener, TcpSocket, TcpStream}; +use tokio::net::{TcpListener, TcpStream}; +use crate::channel::socket::{create_tcp, LocalInterface}; +use crate::ip_proxy::ProxyHandler; use packet::ip::ipv4::packet::IpV4Packet; use packet::tcp::tcp::TcpPacket; -use crate::ip_proxy::ProxyHandler; - #[derive(Clone)] pub struct TcpProxy { port: u16, @@ -19,7 +19,7 @@ pub struct TcpProxy { } impl TcpProxy { - pub async fn new() -> anyhow::Result { + pub async fn new(default_interface: LocalInterface) -> 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)) @@ -28,7 +28,7 @@ impl TcpProxy { let port = tcp_listener.local_addr()?.port(); { let nat_map = nat_map.clone(); - tokio::spawn(tcp_proxy(tcp_listener, nat_map)); + tokio::spawn(tcp_proxy(tcp_listener, nat_map, default_interface)); } Ok(Self { port, nat_map }) } @@ -79,26 +79,33 @@ impl ProxyHandler for TcpProxy { async fn tcp_proxy( tcp_listener: TcpListener, nat_map: Arc>>, + default_interface: LocalInterface, ) { loop { 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() { + let default_interface = default_interface.clone(); 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; - } - }; + let peer_tcp_stream = match tcp_connect( + sender_addr.port(), + dest_addr.into(), + &default_interface, + ) + .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 { @@ -114,15 +121,19 @@ async fn tcp_proxy( } } /// 优先使用来源端口建立tcp连接 -async fn tcp_connect(src_port: u16, addr: SocketAddr) -> anyhow::Result { - let socket = TcpSocket::new_v4()?; +async fn tcp_connect( + src_port: u16, + addr: SocketAddr, + default_interface: &LocalInterface, +) -> anyhow::Result { + let socket = create_tcp(true, default_interface)?; if socket .bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, src_port).into()) .is_err() { socket.bind(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0).into())?; } - let _ = socket.set_nodelay(false); + let _ = socket.set_nodelay(true); let tcp_stream = tokio::time::timeout(Duration::from_secs(5), socket.connect(addr)) .await .with_context(|| format!("TCP connection timeout {}", addr))? diff --git a/vnt/src/ip_proxy/udp_proxy.rs b/vnt/src/ip_proxy/udp_proxy.rs index 0837749..591330a 100644 --- a/vnt/src/ip_proxy/udp_proxy.rs +++ b/vnt/src/ip_proxy/udp_proxy.rs @@ -8,11 +8,11 @@ use std::{collections::HashMap, io, net::SocketAddr}; use parking_lot::Mutex; use tokio::net::UdpSocket; +use crate::channel::socket::{bind_udp, LocalInterface}; +use crate::ip_proxy::ProxyHandler; use packet::ip::ipv4::packet::IpV4Packet; use packet::udp::udp::UdpPacket; -use crate::ip_proxy::ProxyHandler; - #[derive(Clone)] pub struct UdpProxy { port: u16, @@ -20,7 +20,7 @@ pub struct UdpProxy { } impl UdpProxy { - pub async fn new() -> anyhow::Result { + pub async fn new(default_interface: LocalInterface) -> anyhow::Result { let nat_map: Arc>> = Arc::new(Mutex::new(HashMap::with_capacity(16))); let udp = UdpSocket::bind(format!("0.0.0.0:{}", 0)) @@ -29,8 +29,8 @@ impl UdpProxy { let port = udp.local_addr()?.port(); { let nat_map = nat_map.clone(); - tokio::spawn(async { - if let Err(e) = udp_proxy(udp, nat_map).await { + tokio::spawn(async move { + if let Err(e) = udp_proxy(udp, nat_map, default_interface).await { log::warn!("udp_proxy:{:?}", e); } }); @@ -84,7 +84,8 @@ impl ProxyHandler for UdpProxy { async fn udp_proxy( udp: UdpSocket, nat_map: Arc>>, -) -> io::Result<()> { + default_interface: LocalInterface, +) -> anyhow::Result<()> { let mut buf = [0u8; 65536]; let inner_map: Arc, Arc>)>>> = @@ -94,9 +95,15 @@ async fn udp_proxy( 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 + if let Err(e) = udp_proxy0( + &buf[..len], + sender_addr, + &inner_map, + &nat_map, + &udp_socket, + &default_interface, + ) + .await { log::warn!("udp proxy {} {:?}", sender_addr, e); } @@ -116,7 +123,8 @@ async fn udp_proxy0( inner_map: &Arc, Arc>)>>>, map: &Arc>>, udp_socket: &Arc, -) -> io::Result<()> { + default_interface: &LocalInterface, +) -> anyhow::Result<()> { let option = inner_map.lock().get(&sender_addr).cloned(); if let Some((udp, time)) = option { time.store(Instant::now()); @@ -125,11 +133,14 @@ async fn udp_proxy0( 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?, - }; + let peer_udp_socket = match bind_udp( + format!("0.0.0.0:{}", sender_addr.port()).parse().unwrap(), + default_interface, + ) { + Ok(udp) => udp, + Err(_) => bind_udp("0.0.0.0:0".parse().unwrap(), default_interface)?, + }; + let peer_udp_socket = UdpSocket::from_std(peer_udp_socket.into())?; peer_udp_socket.connect(dest_addr).await?; peer_udp_socket.send(buf).await?; let peer_udp_socket = Arc::new(peer_udp_socket); diff --git a/vnt/src/nat/mod.rs b/vnt/src/nat/mod.rs index 4f34810..1ef7b2b 100644 --- a/vnt/src/nat/mod.rs +++ b/vnt/src/nat/mod.rs @@ -11,6 +11,7 @@ use rand::prelude::SliceRandom; use rand::Rng; use crate::channel::punch::{NatInfo, NatType}; +use crate::channel::socket::LocalInterface; use crate::proto::message::PunchNatType; #[cfg(feature = "upnp")] use crate::util::UPnP; @@ -116,6 +117,7 @@ pub struct NatTest { tcp_port: u16, #[cfg(feature = "upnp")] upnp: UPnP, + pub(crate) update_local_ipv4: bool, } impl From for PunchNatType { @@ -144,6 +146,7 @@ impl NatTest { ipv6: Option, udp_ports: Vec, tcp_port: u16, + update_local_ipv4: bool, ) -> NatTest { let ports = vec![0; udp_ports.len()]; let nat_info = NatInfo::new( @@ -178,6 +181,7 @@ impl NatTest { tcp_port, #[cfg(feature = "upnp")] upnp, + update_local_ipv4, } } pub fn can_update(&self) -> bool { @@ -257,6 +261,7 @@ impl NatTest { &self, local_ipv4: Option, ipv6: Option, + default_interface: &LocalInterface, ) -> anyhow::Result { let mut stun_server = self.stun_server.clone(); if stun_server.len() > 5 { @@ -264,7 +269,8 @@ impl NatTest { stun_server.truncate(5); log::info!("stun_server truncate {:?}", stun_server); } - let (nat_type, public_ips, port_range) = stun::stun_test_nat(stun_server)?; + let (nat_type, public_ips, port_range) = + stun::stun_test_nat(stun_server, default_interface)?; if public_ips.is_empty() { Err(anyhow!("public_ips.is_empty"))? } @@ -272,7 +278,9 @@ impl NatTest { guard.nat_type = nat_type; guard.public_ips = public_ips; guard.public_port_range = port_range; - guard.local_ipv4 = local_ipv4; + if local_ipv4.is_some() { + guard.local_ipv4 = local_ipv4; + } guard.ipv6 = ipv6; Ok(guard.clone()) diff --git a/vnt/src/nat/stun.rs b/vnt/src/nat/stun.rs index dcc0e7e..6fcdb9b 100644 --- a/vnt/src/nat/stun.rs +++ b/vnt/src/nat/stun.rs @@ -4,17 +4,21 @@ use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; use std::time::Duration; use crate::channel::punch::NatType; +use crate::channel::socket::{bind_udp, LocalInterface}; use rand::RngCore; use std::net::UdpSocket; use stun_format::Attr; -pub fn stun_test_nat(stun_servers: Vec) -> io::Result<(NatType, Vec, u16)> { +pub fn stun_test_nat( + stun_servers: Vec, + default_interface: &LocalInterface, +) -> anyhow::Result<(NatType, Vec, u16)> { let mut nat_type = NatType::Cone; let mut port_range = 0; let mut hash_set = HashSet::new(); for _ in 0..2 { let stun_servers = stun_servers.clone(); - match stun_test_nat0(stun_servers) { + match stun_test_nat0(stun_servers, default_interface) { Ok((nat_type_t, ip_list_t, port_range_t)) => { if nat_type_t == NatType::Symmetric { nat_type = NatType::Symmetric; @@ -34,8 +38,13 @@ pub fn stun_test_nat(stun_servers: Vec) -> io::Result<(NatType, Vec) -> io::Result<(NatType, Vec, u16)> { - let udp = UdpSocket::bind("0.0.0.0:0")?; +pub fn stun_test_nat0( + stun_servers: Vec, + default_interface: &LocalInterface, +) -> anyhow::Result<(NatType, Vec, u16)> { + let udp = bind_udp("0.0.0.0:0".parse().unwrap(), default_interface)?; + udp.set_nonblocking(false)?; + let udp: UdpSocket = udp.into(); udp.set_read_timeout(Some(Duration::from_millis(500)))?; let mut nat_type = NatType::Cone; let mut min_port = u16::MAX; diff --git a/vnt/src/port_mapping/tcp_mapping.rs b/vnt/src/port_mapping/tcp_mapping.rs index 6bdc61f..82e780c 100644 --- a/vnt/src/port_mapping/tcp_mapping.rs +++ b/vnt/src/port_mapping/tcp_mapping.rs @@ -32,6 +32,7 @@ async fn tcp_mapping_( } async fn copy(source_tcp: TcpStream, destination: &String) -> anyhow::Result<()> { + // 或许这里也应该绑定最匹配的网卡,不然全局代理会影响映射 let dest_tcp = TcpStream::connect(destination) .await .with_context(|| format!("TCP connection target failed {:?}", destination))?; diff --git a/vnt/tun/src/linux/route.rs b/vnt/tun/src/linux/route.rs index b50133c..b6c1630 100644 --- a/vnt/tun/src/linux/route.rs +++ b/vnt/tun/src/linux/route.rs @@ -4,7 +4,16 @@ use std::net::Ipv4Addr; use crate::unix::exe_cmd; pub fn add_route(name: &str, address: Ipv4Addr, netmask: Ipv4Addr) -> io::Result<()> { - let cmd = format!("ip route add {:?}/{:?} dev {}", address, netmask, name); + let cmd = if netmask.is_broadcast() { + format!("route add -host {:?} {}", address, name) + } else { + format!( + "route add -net {}/{} {}", + address, + u32::from(netmask).count_ones(), + name + ) + }; exe_cmd(&cmd)?; Ok(()) }