From e64e17267d202c2c1a032fa7b2069618126acc46 Mon Sep 17 00:00:00 2001 From: lubeilin <1791778603@qq.com> Date: Sun, 7 Jan 2024 21:09:26 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0tcp=E9=80=9A=E9=81=93?= =?UTF-8?q?=E5=A4=84=E7=90=86=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- vnt-cli/src/command/mod.rs | 30 ++-- vnt-cli/src/console_out/mod.rs | 2 +- vnt/src/channel/channel.rs | 251 +++++++++++++++++++++++--------- vnt/src/channel/mod.rs | 2 +- vnt/src/channel/punch.rs | 163 ++++++++++++++++++--- vnt/src/core/mod.rs | 44 ++++-- vnt/src/handle/punch_handler.rs | 11 +- vnt/src/handle/recv_handler.rs | 59 ++++---- vnt/src/nat/mod.rs | 80 +++++----- 9 files changed, 456 insertions(+), 186 deletions(-) diff --git a/vnt-cli/src/command/mod.rs b/vnt-cli/src/command/mod.rs index dbd1759..5acede1 100644 --- a/vnt-cli/src/command/mod.rs +++ b/vnt-cli/src/command/mod.rs @@ -87,8 +87,14 @@ pub fn command_list(vnt: &Vnt) -> Vec { let public_ips: Vec = nat_info.public_ips.iter().map(|v| v.to_string()).collect(); let public_ips = public_ips.join(","); - let local_ip = nat_info.local_ipv4_addr.ip().to_string(); - let ipv6 = nat_info.ipv6_addr.ip().to_string(); + let local_ip = nat_info + .local_ipv4() + .map(|v| v.to_string()) + .unwrap_or("None".to_string()); + let ipv6 = nat_info + .ipv6() + .map(|v| v.to_string()) + .unwrap_or("None".to_string()); (nat_type, public_ips, local_ip, ipv6) } else { ( @@ -100,7 +106,11 @@ pub fn command_list(vnt: &Vnt) -> Vec { }; let (nat_traversal_type, rt) = if let Some(route) = vnt.route(&peer.virtual_ip) { let nat_traversal_type = if route.metric == 1 { - "p2p" + if route.is_tcp { + "tcp-p2p" + } else { + "p2p" + } } else if route.addr == info.connect_server { "server-relay" } else { @@ -148,12 +158,14 @@ pub fn command_info(vnt: &Vnt) -> Info { let nat_type = format!("{:?}", nat_info.nat_type); let public_ips: Vec = nat_info.public_ips.iter().map(|v| v.to_string()).collect(); let public_ips = public_ips.join(","); - let local_addr = nat_info.local_ipv4_addr.to_string(); - let ipv6_addr = if nat_info.ipv6_addr.ip().is_unspecified() { - "None".to_string() - } else { - nat_info.ipv6_addr.ip().to_string() - }; + let local_addr = nat_info + .local_ipv4() + .map(|v| v.to_string()) + .unwrap_or("None".to_string()); + let ipv6_addr = nat_info + .ipv6() + .map(|v| v.to_string()) + .unwrap_or("None".to_string()); Info { name, virtual_ip, diff --git a/vnt-cli/src/console_out/mod.rs b/vnt-cli/src/console_out/mod.rs index 8bb746f..ce88246 100644 --- a/vnt-cli/src/console_out/mod.rs +++ b/vnt-cli/src/console_out/mod.rs @@ -76,7 +76,7 @@ pub fn console_device_list(mut list: Vec) { ("".to_string(), Style::new().red()), ]); } else { - if &item.nat_traversal_type == "p2p" { + if item.nat_traversal_type.contains("p2p") { out_list.push(vec![ (item.name, Style::new().green()), (item.virtual_ip, Style::new().green()), diff --git a/vnt/src/channel/channel.rs b/vnt/src/channel/channel.rs index 792c20c..fde3e89 100644 --- a/vnt/src/channel/channel.rs +++ b/vnt/src/channel/channel.rs @@ -1,9 +1,13 @@ use std::collections::HashMap; use std::io::{Read, Write}; -use std::net::UdpSocket as StdUdpSocket; -use std::net::{Ipv4Addr, Shutdown, SocketAddr}; +use std::net::{Ipv4Addr, Ipv6Addr, Shutdown, SocketAddr}; use std::net::{SocketAddrV6, TcpStream}; -use std::sync::atomic::{AtomicBool, Ordering}; +use std::net::{TcpListener, UdpSocket as StdUdpSocket}; +#[cfg(any(unix))] +use std::os::fd::AsRawFd; +#[cfg(target_os = "windows")] +use std::os::windows::io::AsRawSocket; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::{Duration, Instant}; use std::{io, thread}; @@ -33,6 +37,7 @@ pub struct ContextInner { current_device: Arc>, first_latency: bool, is_close: AtomicBool, + tcp_port: u16, } #[derive(Clone)] @@ -47,6 +52,7 @@ impl Context { current_device: Arc>, _channel_num: usize, first_latency: bool, + tcp_port: u16, ) -> Self { //当前版本只支持一个通道 let channel_num = 1; @@ -64,6 +70,7 @@ impl Context { current_device, first_latency, is_close: AtomicBool::new(false), + tcp_port, }); Self { inner } } @@ -77,6 +84,7 @@ impl Context { *self.inner.status_receiver.borrow() == Status::Cone } pub fn close(&self) -> io::Result<()> { + let last = self.is_close(); self.inner.is_close.store(true, Ordering::Release); let _ = self.inner.status_sender.send(Status::Close); if let Ok(port) = self.main_local_udp_port() { @@ -99,9 +107,22 @@ impl Context { log::error!("发送停止消息到tcp失败:{:?}", e); } } - for (_, tcp) in self.inner.tcp_map.read().clone() { - if let Err(e) = tcp.lock().shutdown(Shutdown::Both) { - log::error!("发送停止消息到tcp失败:{:?}", e); + if !last { + for (_, tcp) in self.inner.tcp_map.read().clone() { + if let Err(e) = tcp.lock().shutdown(Shutdown::Both) { + log::error!("发送停止消息到tcp失败:{:?}", e); + } + } + if let Err(e) = TcpStream::connect_timeout( + &SocketAddr::V6(SocketAddrV6::new( + Ipv6Addr::LOCALHOST, + self.inner.tcp_port, + 0, + 0, + )), + Duration::from_secs(1), + ) { + log::error!("发送停止消息到tcp_listener失败:{:?}", e); } } Ok(()) @@ -152,18 +173,15 @@ impl Context { #[inline] pub fn send_main_tcp(&self, buf: &[u8]) -> io::Result { if let Some(sender) = &self.inner.main_tcp_channel { - let mut stream = sender.lock(); - let mut head = [0; 4]; - let len = buf.len(); - head[2] = (len >> 8) as u8; - head[3] = (len & 0xFF) as u8; - stream.write_all(&head)?; - stream.write_all(buf)?; - Ok(len) + Self::send_tcp(sender, buf) } else { return Err(io::Error::new(io::ErrorKind::NotFound, "tcp not found")); } } + pub fn send_tcp(sender: &Mutex, buf: &[u8]) -> io::Result { + let mut stream = sender.lock(); + send_tcp(&mut stream, buf) + } pub fn send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result { if let Some(sender) = &self.inner.main_tcp_channel { @@ -229,8 +247,14 @@ impl Context { TCP_ID => self.send_main_tcp(buf), UDP_ID => self.send_main_udp(buf, route_key.addr), _ => { - if let Some(udp) = self.get_udp_by_route(route_key) { - return udp.send_to(buf, route_key.addr).await; + if route_key.is_tcp { + if let Some(tcp) = self.get_tcp_by_route(route_key) { + return Self::send_tcp(&tcp, buf); + } + } else { + if let Some(udp) = self.get_udp_by_route(route_key) { + return udp.send_to(buf, route_key.addr).await; + } } Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) } @@ -241,16 +265,27 @@ impl Context { TCP_ID => self.send_main_tcp(buf), UDP_ID => self.send_main_udp(buf, route_key.addr), _ => { - if let Some(udp) = self.get_udp_by_route(route_key) { - return udp.try_send_to(buf, route_key.addr); + if route_key.is_tcp { + if let Some(tcp) = self.get_tcp_by_route(route_key) { + return Self::send_tcp(&tcp, buf); + } + } else { + if let Some(udp) = self.get_udp_by_route(route_key) { + return udp.try_send_to(buf, route_key.addr); + } } Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) } } } + #[inline] fn get_udp_by_route(&self, route_key: &RouteKey) -> Option> { self.inner.udp_map.read().get(&route_key.index).cloned() } + #[inline] + fn get_tcp_by_route(&self, route_key: &RouteKey) -> Option>> { + self.inner.tcp_map.read().get(&route_key.index).cloned() + } pub fn add_route_if_absent(&self, id: Ipv4Addr, route: Route) { self.add_route_(id, route, true) @@ -385,59 +420,33 @@ impl Context { pub struct Channel { context: Context, handler: ChannelDataHandler, + tcp_listener: TcpListener, } impl Channel { - pub fn new(context: Context, handler: ChannelDataHandler) -> Self { - Self { context, handler } - } -} - -impl Channel { - fn tcp_handle( - tcp_r: &mut TcpStream, - context: &Context, - handler: &ChannelDataHandler, - head_reserve: usize, - ) -> io::Result<()> { - let mut head = [0; 4]; - let addr = tcp_r.peer_addr()?; - let key = RouteKey::new(true, TCP_ID, addr); - loop { - if context.is_close() { - return Ok(()); - } - let mut buf = [0; 4096]; - tcp_r.read_exact(&mut head)?; - let len = (((head[2] as u16) << 8) | head[3] as u16) as usize; - if len < 12 || len > buf.len() { - return Err(io::Error::new( - io::ErrorKind::InvalidData, - "length overflow", - )); - } - tcp_r.read_exact(&mut buf[head_reserve..head_reserve + len])?; - handler.handle(&mut buf, head_reserve, head_reserve + len, key, context); + pub fn new(context: Context, handler: ChannelDataHandler, tcp_listener: TcpListener) -> Self { + Self { + context, + handler, + tcp_listener, } } - fn start_tcp( - mut tcp_stream: TcpStream, - context: Context, - handler: ChannelDataHandler, - head_reserve: usize, - ) { +} + +impl Channel { + fn start_tcp(mut tcp_stream: TcpStream, context: Context, handler: ChannelDataHandler) { let current_device = context.inner.current_device.clone(); loop { if let Err(e) = tcp_stream.set_nodelay(true) { log::info!("set_nodelay:{:?}", e); } - if let Err(e) = tcp_stream.set_write_timeout(Some(Duration::from_secs(3))) { + if let Err(e) = tcp_stream.set_write_timeout(Some(Duration::from_secs(5))) { log::info!("set_write_timeout:{:?}", e); } if let Err(e) = tcp_stream.set_read_timeout(Some(Duration::from_secs(10))) { log::info!("set_read_timeout:{:?}", e); } - if let Err(e) = Self::tcp_handle(&mut tcp_stream, &context, &handler, head_reserve) { + if let Err(e) = tcp_handle(TCP_ID, &mut tcp_stream, &context, &handler) { log::info!("tcp链接断开:{:?}", e); } if let Err(e) = tcp_stream.shutdown(Shutdown::Both) { @@ -463,12 +472,50 @@ impl Channel { } } } + fn start_tcp_listen( + worker: VntWorker, + context: Context, + handler: ChannelDataHandler, + tcp_listener: TcpListener, + ) { + let counter = Arc::new(AtomicUsize::new(0)); + for stream in tcp_listener.incoming() { + if context.is_close() { + break; + } + if counter.load(Ordering::Relaxed) > 20 { + continue; + } + match stream { + Ok(stream) => { + let context = context.clone(); + let handler = handler.clone(); + let counter = counter.clone(); + counter.fetch_add(1, Ordering::Relaxed); + thread::spawn(move || { + if let Err(e) = start_tcp_handle(stream, context, handler) { + log::error!("{:?}", e); + } + counter.fetch_sub(1, Ordering::Relaxed); + }); + } + Err(e) => { + log::error!("connection failed {:?}", e); + } + } + } + for (_, tcp) in context.inner.tcp_map.read().clone() { + if let Err(e) = tcp.lock().shutdown(Shutdown::Both) { + log::error!("发送停止消息到tcp失败:{:?}", e); + } + } + worker.stop_all(); + } pub async fn start( self, mut worker: VntWorker, tcp: Option, - head_reserve: usize, //头部预留字节 symmetric_channel_num: usize, //对称网络,则再加一组监听,提升打洞成功率 relay: bool, ) { @@ -482,7 +529,7 @@ impl Channel { thread::Builder::new() .name("channel_tcp".into()) .spawn(move || { - Self::start_tcp(tcp_stream, context, handler, head_reserve); + Self::start_tcp(tcp_stream, context, handler); drop(main_channel_tcp) }) .unwrap(); @@ -496,7 +543,7 @@ impl Channel { .name("channel_udp".into()) .spawn(move || { log::info!("启动udp v4"); - Self::main_start_(worker, context, UDP_ID, main_channel, handler, head_reserve) + Self::main_start_(worker, context, UDP_ID, main_channel, handler) }) .unwrap(); } @@ -504,6 +551,19 @@ impl Channel { worker.stop_wait().await; return; } + { + let context = context.clone(); + let handler = handler.clone(); + let tcp_listener = self.tcp_listener; + let worker = worker.worker("tcp_listener"); + thread::Builder::new() + .name("tcp_listener".into()) + .spawn(move || { + log::info!("启动tcp"); + Self::start_tcp_listen(worker, context, handler, tcp_listener) + }) + .unwrap(); + } let mut cur_status = Status::Cone; let mut status_receiver = context.inner.status_receiver.clone(); let channel_num = context.inner.channel_num; @@ -530,7 +590,7 @@ impl Channel { Ok(udp) => { let udp = Arc::new(udp); let context = context.clone(); - tokio::spawn(Self::start_(worker.worker("symmetric_channel"),context, udp,handler.clone(), head_reserve)); + tokio::spawn(Self::start_(worker.worker("symmetric_channel"),context, udp,handler.clone())); } Err(e) => { log::error!("{}",e); @@ -558,9 +618,9 @@ impl Channel { id: usize, udp: StdUdpSocket, handler: ChannelDataHandler, - head_reserve: usize, ) { let mut buf = [0; 4096]; + let head_reserve = handler.head_reserve; loop { match udp.recv_from(&mut buf[head_reserve..]) { Ok((len, addr)) => { @@ -591,20 +651,17 @@ impl Channel { context: Context, udp: Arc, handler: ChannelDataHandler, - head_reserve: usize, ) { let mut status_receiver = context.inner.status_receiver.clone(); - #[cfg(target_os = "windows")] - use std::os::windows::io::AsRawSocket; + #[cfg(target_os = "windows")] let id = 3 + udp.as_raw_socket() as usize; #[cfg(any(unix))] - use std::os::fd::AsRawFd; - #[cfg(any(unix))] let id = 3 + udp.as_raw_fd() as usize; context.insert_udp(id, udp.clone()); let mut buf = [0; 4096]; + let head_reserve = handler.head_reserve; loop { tokio::select! { rs=udp.recv_from(&mut buf[head_reserve..])=>{ @@ -643,3 +700,65 @@ impl Channel { context.remove_udp(id); } } +pub fn start_tcp_handle( + mut stream: TcpStream, + context: Context, + handler: ChannelDataHandler, +) -> io::Result<()> { + stream.set_write_timeout(Some(Duration::from_secs(5)))?; + stream.set_read_timeout(Some(Duration::from_secs(10)))?; + if let Err(e) = stream.set_nodelay(true) { + log::error!("设置nodelay失败 {:?}", e); + } + let writer = stream.try_clone()?; + #[cfg(target_os = "windows")] + let id = 3 + stream.as_raw_socket() as usize; + #[cfg(any(unix))] + let id = 3 + stream.as_raw_fd() as usize; + context + .inner + .tcp_map + .write() + .insert(id, Arc::new(Mutex::new(writer))); + if let Err(e) = tcp_handle(id, &mut stream, &context, &handler) { + log::error!("tcp_handle {:?}", e); + } + context.inner.tcp_map.write().remove(&id); + Ok(()) +} +pub fn tcp_handle( + id: usize, + tcp_r: &mut TcpStream, + context: &Context, + handler: &ChannelDataHandler, +) -> io::Result<()> { + let mut head = [0; 4]; + let addr = tcp_r.peer_addr()?; + let key = RouteKey::new(true, id, addr); + let head_reserve = handler.head_reserve; + loop { + if context.is_close() { + return Ok(()); + } + let mut buf = [0; 4096]; + tcp_r.read_exact(&mut head)?; + let len = (((head[2] as u16) << 8) | head[3] as u16) as usize; + if len < 12 || len > buf.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "length overflow", + )); + } + tcp_r.read_exact(&mut buf[head_reserve..head_reserve + len])?; + handler.handle(&mut buf, head_reserve, head_reserve + len, key, context); + } +} +pub fn send_tcp(stream: &mut TcpStream, buf: &[u8]) -> io::Result { + let mut head = [0; 4]; + let len = buf.len(); + head[2] = (len >> 8) as u8; + head[3] = (len & 0xFF) as u8; + stream.write_all(&head)?; + stream.write_all(buf)?; + Ok(len) +} diff --git a/vnt/src/channel/mod.rs b/vnt/src/channel/mod.rs index 7bdca0c..02d0267 100644 --- a/vnt/src/channel/mod.rs +++ b/vnt/src/channel/mod.rs @@ -17,7 +17,7 @@ pub enum Status { #[derive(Copy, Clone, Debug)] pub struct Route { - is_tcp: bool, + pub is_tcp: bool, index: usize, pub addr: SocketAddr, pub metric: u8, diff --git a/vnt/src/channel/punch.rs b/vnt/src/channel/punch.rs index 9dafbf1..71bc121 100644 --- a/vnt/src/channel/punch.rs +++ b/vnt/src/channel/punch.rs @@ -1,12 +1,13 @@ use std::collections::HashMap; -use std::io; -use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; +use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6, TcpStream}; use std::str::FromStr; use std::time::Duration; +use std::{io, thread}; use rand::prelude::SliceRandom; -use crate::channel::channel::Context; +use crate::channel::channel::{send_tcp, start_tcp_handle, Context}; +use crate::handle::recv_handler::ChannelDataHandler; #[derive(Copy, Clone, Eq, PartialEq, Debug)] pub enum PunchModel { @@ -32,9 +33,11 @@ pub struct NatInfo { pub public_ips: Vec, pub public_port: u16, pub public_port_range: u16, - pub local_ipv4_addr: SocketAddrV4, - pub ipv6_addr: SocketAddrV6, pub nat_type: NatType, + pub(crate) local_ipv4: Option, + pub(crate) ipv6: Option, + pub(crate) udp_port: u16, + pub tcp_port: u16, } #[derive(Clone, Copy, PartialEq, Eq, Debug, Hash)] @@ -48,8 +51,10 @@ impl NatInfo { mut public_ips: Vec, public_port: u16, public_port_range: u16, - local_ipv4_addr: SocketAddrV4, - ipv6_addr: SocketAddrV6, + mut local_ipv4: Option, + mut ipv6: Option, + udp_port: u16, + tcp_port: u16, mut nat_type: NatType, ) -> Self { public_ips.retain(|ip| { @@ -62,12 +67,24 @@ impl NatInfo { if public_ips.len() > 1 { nat_type = NatType::Symmetric; } + if let Some(ip) = local_ipv4 { + if ip.is_multicast() || ip.is_broadcast() || ip.is_unspecified() || ip.is_loopback() { + local_ipv4 = None + } + } + if let Some(ip) = ipv6 { + if ip.is_multicast() || ip.is_unspecified() || ip.is_loopback() { + ipv6 = None + } + } Self { public_ips, public_port, public_port_range, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, nat_type, } } @@ -85,6 +102,53 @@ impl NatInfo { } } } + pub fn local_ipv4(&self) -> Option { + self.local_ipv4 + } + pub fn ipv6(&self) -> Option { + self.ipv6 + } + pub fn local_udp_ipv4addr(&self) -> Option { + if self.udp_port == 0 { + return None; + } + if let Some(local_ipv4) = self.local_ipv4 { + Some(SocketAddr::V4(SocketAddrV4::new(local_ipv4, self.udp_port))) + } else { + None + } + } + pub fn local_udp_ipv6addr(&self) -> Option { + if self.udp_port == 0 { + return None; + } + if let Some(ipv6) = self.ipv6 { + Some(SocketAddr::V6(SocketAddrV6::new(ipv6, self.udp_port, 0, 0))) + } else { + None + } + } + + pub fn local_tcp_ipv6addr(&self) -> Option { + if self.tcp_port == 0 { + return None; + } + if let Some(ipv6) = self.ipv6 { + Some(SocketAddr::V6(SocketAddrV6::new(ipv6, self.tcp_port, 0, 0))) + } else { + None + } + } + pub fn local_tcp_ipv4addr(&self) -> Option { + if self.tcp_port == 0 { + return None; + } + if let Some(ipv4) = self.local_ipv4 { + Some(SocketAddr::V4(SocketAddrV4::new(ipv4, self.tcp_port))) + } else { + None + } + } } #[derive(Clone)] @@ -93,10 +157,17 @@ pub struct Punch { port_vec: Vec, port_index: HashMap, punch_model: PunchModel, + is_tcp: bool, + handler: ChannelDataHandler, } impl Punch { - pub fn new(context: Context, punch_model: PunchModel) -> Self { + pub fn new( + context: Context, + punch_model: PunchModel, + is_tcp: bool, + handler: ChannelDataHandler, + ) -> Self { let mut port_vec: Vec = (1..65535).collect(); port_vec.push(65535); let mut rng = rand::thread_rng(); @@ -106,30 +177,74 @@ impl Punch { port_vec, port_index: HashMap::new(), punch_model, + is_tcp, + handler, } } } impl Punch { + fn connect_tcp(&self, buf: &[u8], addr: &SocketAddr) -> bool { + match TcpStream::connect_timeout(&addr, Duration::from_secs(1)) { + Ok(mut tcp_stream) => { + let context = self.context.clone(); + let handler = self.handler.clone(); + match send_tcp(&mut tcp_stream, buf) { + Ok(_) => {} + Err(e) => { + log::warn!("发送到tcp失败,addr={},err={}", addr, e); + return false; + } + } + thread::spawn(move || { + if let Err(e) = start_tcp_handle(tcp_stream, context, handler) { + log::error!("{:?}", e); + } + }); + return true; + } + Err(e) => { + log::warn!("连接到tcp失败,addr={},err={}", addr, e); + } + } + false + } pub async fn punch(&mut self, buf: &[u8], id: Ipv4Addr, nat_info: NatInfo) -> io::Result<()> { if !self.context.need_punch(&id) { return Ok(()); } - if !nat_info.local_ipv4_addr.ip().is_unspecified() && nat_info.local_ipv4_addr.port() != 0 { - let _ = self - .context - .send_main_udp(buf, SocketAddr::V4(nat_info.local_ipv4_addr)); + if self.is_tcp { + //向tcp发起连接 + if let Some(ipv6_addr) = nat_info.local_tcp_ipv6addr() { + if self.connect_tcp(buf, &ipv6_addr) { + return Ok(()); + } + } + log::info!("local_tcp_ipv4addr={:?}", nat_info.local_tcp_ipv4addr()); + //向tcp发起连接 + if let Some(ipv4_addr) = nat_info.local_tcp_ipv4addr() { + if self.connect_tcp(buf, &ipv4_addr) { + return Ok(()); + } + } + if nat_info.nat_type == NatType::Cone && nat_info.public_ips.len() == 1 { + let addr = + SocketAddr::V4(SocketAddrV4::new(nat_info.public_ips[0], nat_info.tcp_port)); + if self.connect_tcp(buf, &addr) { + return Ok(()); + } + } } - if self.punch_model != PunchModel::IPv4 - && !nat_info.ipv6_addr.ip().is_unspecified() - && nat_info.ipv6_addr.port() != 0 - { - let rs = self - .context - .send_main_udp(buf, SocketAddr::V6(nat_info.ipv6_addr)); - log::info!("发送到ipv6地址:{:?},rs={:?}", nat_info.ipv6_addr, rs); - if rs.is_ok() && self.punch_model == PunchModel::IPv6 { - return Ok(()); + if let Some(ipv4_addr) = nat_info.local_udp_ipv4addr() { + let _ = self.context.send_main_udp(buf, ipv4_addr); + } + if self.punch_model != PunchModel::IPv4 { + if let Some(ipv6_addr) = nat_info.local_udp_ipv6addr() { + let rs = self.context.send_main_udp(buf, ipv6_addr); + log::info!("发送到ipv6地址:{:?},rs={:?}", ipv6_addr, rs); + if rs.is_ok() && self.punch_model == PunchModel::IPv6 { + return Ok(()); + } } } match nat_info.nat_type { diff --git a/vnt/src/core/mod.rs b/vnt/src/core/mod.rs index c07c4d1..644dd43 100644 --- a/vnt/src/core/mod.rs +++ b/vnt/src/core/mod.rs @@ -1,8 +1,8 @@ use std::collections::HashMap; use std::io; -use std::net::TcpStream; use std::net::UdpSocket; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; +use std::net::{TcpListener, TcpStream}; use std::sync::Arc; use std::time::Duration; @@ -241,14 +241,16 @@ impl VntUtil { } else { (None, None) }; + let tcp_listener = TcpListener::bind(format!("[::]:{}", config.port))?; + let local_tcp_port = tcp_listener.local_addr()?.port(); let context = Context::new( self.main_channel, tcp_sender, current_device.clone(), 1, config.first_latency, + local_tcp_port, ); - let punch = Punch::new(context.clone(), config.punch_model); let idle = Idle::new(Duration::from_secs(16), context.clone()); let channel_sender = ChannelSender::new(context.clone()); @@ -268,17 +270,19 @@ impl VntUtil { let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected)); let public_ip = response.public_ip; let public_port = response.public_port; - let local_port = context.main_local_udp_port().unwrap_or(0); + let local_udp_port = context.main_local_udp_port().unwrap_or(0); + let local_ipv4 = crate::nat::local_ipv4(); + let ipv6 = crate::nat::local_ipv6(); - let local_ipv4_addr = crate::nat::local_ipv4_addr(local_port); - let ipv6_addr = crate::nat::local_ipv6_addr(local_port); // NAT检测 let nat_test = NatTest::new( config.stun_server.clone(), public_ip, public_port, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + local_udp_port, + local_tcp_port, ); let in_external_route = if config.in_ips.is_empty() { None @@ -375,16 +379,21 @@ impl VntUtil { self.rsa_cipher.clone(), config.relay, config.token.clone(), + 14, + ); + let punch = Punch::new( + context.clone(), + config.punch_model, + config.tcp, + channel_recv_handler.clone(), ); { - let channel = Channel::new(context.clone(), channel_recv_handler); + let channel = Channel::new(context.clone(), channel_recv_handler, tcp_listener); let channel_worker = vnt_status_manager.worker("channel_worker"); let relay = config.relay; - tokio::spawn(async move { - channel - .start(channel_worker, tcp_receiver, 14, 65, relay) - .await - }); + tokio::spawn( + async move { channel.start(channel_worker, tcp_receiver, 65, relay).await }, + ); } { let nat_test = nat_test.clone(); @@ -455,7 +464,14 @@ impl VntUtil { let nat_test = nat_test.clone(); tokio::spawn(async move { let info = nat_test - .re_test(public_ip, public_port, local_ipv4_addr, ipv6_addr) + .re_test( + public_ip, + public_port, + local_ipv4, + ipv6, + local_udp_port, + local_tcp_port, + ) .await; context.switch(info.nat_type); }); diff --git a/vnt/src/handle/punch_handler.rs b/vnt/src/handle/punch_handler.rs index d7aaa0a..269cd4b 100644 --- a/vnt/src/handle/punch_handler.rs +++ b/vnt/src/handle/punch_handler.rs @@ -158,11 +158,12 @@ pub fn punch_packet( .collect(); punch_reply.public_port = nat_info.public_port as u32; punch_reply.public_port_range = nat_info.public_port_range as u32; - punch_reply.local_ip = u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets()); - punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32; - if !nat_info.ipv6_addr.ip().is_unspecified() { - punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32; - punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec(); + punch_reply.local_ip = u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED)); + punch_reply.local_port = nat_info.udp_port as u32; + punch_reply.tcp_port = nat_info.tcp_port as u32; + if let Some(ipv6) = nat_info.ipv6 { + punch_reply.ipv6_port = nat_info.udp_port as u32; + punch_reply.ipv6 = ipv6.octets().to_vec(); } punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type)); let bytes = punch_reply.write_to_bytes()?; diff --git a/vnt/src/handle/recv_handler.rs b/vnt/src/handle/recv_handler.rs index 397fb96..b955643 100644 --- a/vnt/src/handle/recv_handler.rs +++ b/vnt/src/handle/recv_handler.rs @@ -1,5 +1,5 @@ use std::collections::HashMap; -use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6}; +use std::net::{Ipv4Addr, Ipv6Addr}; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -57,6 +57,7 @@ pub struct ChannelDataHandler { relay: bool, token: String, time: Arc>, + pub head_reserve: usize, } impl ChannelDataHandler { @@ -78,6 +79,7 @@ impl ChannelDataHandler { rsa_cipher: Option, relay: bool, token: String, + head_reserve: usize, ) -> Self { Self { current_device, @@ -99,6 +101,7 @@ impl ChannelDataHandler { relay, token, time: Arc::new(AtomicCell::new(Instant::now())), + head_reserve, } } } @@ -400,23 +403,24 @@ impl ChannelDataHandler { .iter() .map(|v| Ipv4Addr::from(v.to_be_bytes())) .collect(); - let local_ipv4_addr = SocketAddrV4::new( - Ipv4Addr::from(punch_info.local_ip.to_be_bytes()), - punch_info.local_port as u16, - ); - let ipv6_addr = if punch_info.ipv6.len() == 16 { + let local_ipv4 = Some(Ipv4Addr::from(punch_info.local_ip.to_be_bytes())); + let udp_port = punch_info.local_port as u16; + let tcp_port = punch_info.tcp_port as u16; + let ipv6 = if punch_info.ipv6.len() == 16 { let ipv6: [u8; 16] = punch_info.ipv6.try_into().unwrap(); - SocketAddrV6::new(Ipv6Addr::from(ipv6), punch_info.ipv6_port as u16, 0, 0) + Some(Ipv6Addr::from(ipv6)) } else { - SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0) + None }; let peer_nat_info = NatInfo::new( public_ips, punch_info.public_port as u16, punch_info.public_port_range as u16, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, punch_info.nat_type.enum_value_or_default().into(), ); { @@ -437,11 +441,11 @@ impl ChannelDataHandler { punch_reply.nat_type = protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type)); punch_reply.local_ip = - u32::from_be_bytes(nat_info.local_ipv4_addr.ip().octets()); - punch_reply.local_port = nat_info.local_ipv4_addr.port() as u32; - if !nat_info.ipv6_addr.ip().is_unspecified() { - punch_reply.ipv6 = nat_info.ipv6_addr.ip().octets().to_vec(); - punch_reply.ipv6_port = nat_info.ipv6_addr.port() as u32; + u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED)); + punch_reply.local_port = nat_info.udp_port as u32; + if let Some(ipv6) = nat_info.ipv6() { + punch_reply.ipv6 = ipv6.octets().to_vec(); + punch_reply.ipv6_port = nat_info.udp_port as u32; } let bytes = punch_reply.write_to_bytes()?; let mut punch_packet = @@ -453,18 +457,6 @@ impl ChannelDataHandler { punch_packet.set_source(current_device.virtual_ip()); punch_packet.set_destination(source); punch_packet.set_payload(&bytes)?; - // if !peer_nat_info.local_ip.is_unspecified() && peer_nat_info.local_port != 0 { - // let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED])?; - // packet.set_version(Version::V1); - // packet.first_set_ttl(1); - // packet.set_protocol(Protocol::Control); - // packet.set_transport_protocol(control_packet::Protocol::PunchRequest.into()); - // packet.set_source(current_device.virtual_ip()); - // packet.set_destination(source); - // self.client_cipher.encrypt_ipv4(&mut packet)?; - // let _ = context.try_send_main_udp(packet.buffer(), - // SocketAddr::V4(SocketAddrV4::new(peer_nat_info.local_ip, peer_nat_info.local_port))); - // } if self.punch(source, peer_nat_info) { self.client_cipher.encrypt_ipv4(&mut punch_packet)?; context.try_send_by_key(punch_packet.buffer(), route_key)?; @@ -592,15 +584,18 @@ impl ChannelDataHandler { .build() .unwrap() .block_on(async move { - let local_port = context.main_local_udp_port().unwrap_or(0); - let local_ipv4_addr = nat::local_ipv4_addr(local_port); - let ipv6_addr = nat::local_ipv6_addr(local_port); + let local_ipv4 = nat::local_ipv4(); + let ipv6 = nat::local_ipv6(); + let udp_port = nat_test.nat_info().udp_port; + let tcp_port = nat_test.nat_info().tcp_port; let nat_info = nat_test .re_test( Ipv4Addr::from(response.public_ip), response.public_port as u16, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, ) .await; context.switch(nat_info.nat_type); diff --git a/vnt/src/nat/mod.rs b/vnt/src/nat/mod.rs index 66cb310..223440d 100644 --- a/vnt/src/nat/mod.rs +++ b/vnt/src/nat/mod.rs @@ -1,10 +1,10 @@ -use crossbeam_utils::atomic::AtomicCell; use std::io; use std::net::UdpSocket; -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddrV4, SocketAddrV6}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; use std::sync::Arc; use std::time::{Duration, Instant}; +use crossbeam_utils::atomic::AtomicCell; use parking_lot::Mutex; use crate::channel::punch::{NatInfo, NatType}; @@ -12,7 +12,7 @@ use crate::proto::message::PunchNatType; mod stun_test; -pub fn local_ipv4() -> io::Result { +pub fn local_ipv4_() -> io::Result { let socket = UdpSocket::bind("0.0.0.0:0")?; socket.connect("8.8.8.8:80")?; let addr = socket.local_addr()?; @@ -21,8 +21,17 @@ pub fn local_ipv4() -> io::Result { IpAddr::V6(_) => Ok(Ipv4Addr::UNSPECIFIED), } } +pub fn local_ipv4() -> Option { + match local_ipv4_() { + Ok(ipv4) => Some(ipv4), + Err(e) => { + log::warn!("获取ipv4失败:{:?}", e); + None + } + } +} -pub fn local_ipv6() -> io::Result { +pub fn local_ipv6_() -> io::Result { let socket = UdpSocket::bind("[::]:0")?; socket.connect("[2001:4860:4860:0000:0000:0000:0000:8888]:80")?; let addr = socket.local_addr()?; @@ -31,23 +40,12 @@ pub fn local_ipv6() -> io::Result { IpAddr::V6(ip) => Ok(ip), } } - -pub fn local_ipv4_addr(port: u16) -> SocketAddrV4 { - match local_ipv4() { - Ok(ipv4) => SocketAddrV4::new(ipv4, port), +pub fn local_ipv6() -> Option { + match local_ipv6_() { + Ok(ipv6) => Some(ipv6), Err(e) => { - log::warn!("获取本地ipv4地址失败:{}", e); - SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0) - } - } -} - -pub fn local_ipv6_addr(port: u16) -> SocketAddrV6 { - match local_ipv6() { - Ok(ipv6) => SocketAddrV6::new(ipv6, port, 0, 0), - Err(e) => { - log::warn!("获取本地ipv6地址失败:{}", e); - SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, 0, 0, 0) + log::warn!("获取ipv6失败:{:?}", e); + None } } } @@ -82,8 +80,10 @@ impl NatTest { mut stun_server: Vec, public_ip: Ipv4Addr, public_port: u16, - local_ipv4_addr: SocketAddrV4, - ipv6_addr: SocketAddrV6, + local_ipv4: Option, + ipv6: Option, + udp_port: u16, + tcp_port: u16, ) -> NatTest { let server = stun_server[0].clone(); stun_server.resize(3, server); @@ -91,8 +91,10 @@ impl NatTest { vec![public_ip], public_port, 0, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, NatType::Cone, ); let info = Arc::new(Mutex::new(nat_info)); @@ -118,15 +120,19 @@ impl NatTest { &self, public_ip: Ipv4Addr, public_port: u16, - local_ipv4_addr: SocketAddrV4, - ipv6_addr: SocketAddrV6, + local_ipv4: Option, + ipv6: Option, + udp_port: u16, + tcp_port: u16, ) -> NatInfo { let info = NatTest::re_test_( &self.stun_server, public_ip, public_port, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, ) .await; log::info!("探测nat类型={:?}", info); @@ -137,8 +143,10 @@ impl NatTest { stun_server: &Vec, public_ip: Ipv4Addr, public_port: u16, - local_ipv4_addr: SocketAddrV4, - ipv6_addr: SocketAddrV6, + local_ipv4: Option, + ipv6: Option, + udp_port: u16, + tcp_port: u16, ) -> NatInfo { return match stun_test::stun_test_nat(stun_server.clone()).await { Ok((nat_type, mut public_ips, port_range)) => { @@ -149,8 +157,10 @@ impl NatTest { public_ips, public_port, port_range, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, nat_type, ) } @@ -160,8 +170,10 @@ impl NatTest { vec![public_ip], public_port, 0, - local_ipv4_addr, - ipv6_addr, + local_ipv4, + ipv6, + udp_port, + tcp_port, NatType::Cone, ) }