增加tcp通道处理逻辑

This commit is contained in:
lubeilin
2024-01-07 21:09:26 +08:00
parent 7098111ad1
commit e64e17267d
9 changed files with 456 additions and 186 deletions
+21 -9
View File
@@ -87,8 +87,14 @@ pub fn command_list(vnt: &Vnt) -> Vec<DeviceItem> {
let public_ips: Vec<String> =
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<DeviceItem> {
};
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<String> = 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,
+1 -1
View File
@@ -76,7 +76,7 @@ pub fn console_device_list(mut list: Vec<DeviceItem>) {
("".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()),
+185 -66
View File
@@ -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<AtomicCell<CurrentDeviceInfo>>,
first_latency: bool,
is_close: AtomicBool,
tcp_port: u16,
}
#[derive(Clone)]
@@ -47,6 +52,7 @@ impl Context {
current_device: Arc<AtomicCell<CurrentDeviceInfo>>,
_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<usize> {
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<TcpStream>, buf: &[u8]) -> io::Result<usize> {
let mut stream = sender.lock();
send_tcp(&mut stream, buf)
}
pub fn send_main(&self, buf: &[u8], addr: SocketAddr) -> io::Result<usize> {
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<Arc<UdpSocket>> {
self.inner.udp_map.read().get(&route_key.index).cloned()
}
#[inline]
fn get_tcp_by_route(&self, route_key: &RouteKey) -> Option<Arc<Mutex<TcpStream>>> {
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<TcpStream>,
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<UdpSocket>,
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<usize> {
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)
}
+1 -1
View File
@@ -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,
+139 -24
View File
@@ -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<Ipv4Addr>,
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<Ipv4Addr>,
pub(crate) ipv6: Option<Ipv6Addr>,
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<Ipv4Addr>,
public_port: u16,
public_port_range: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
mut local_ipv4: Option<Ipv4Addr>,
mut ipv6: Option<Ipv6Addr>,
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<Ipv4Addr> {
self.local_ipv4
}
pub fn ipv6(&self) -> Option<Ipv6Addr> {
self.ipv6
}
pub fn local_udp_ipv4addr(&self) -> Option<SocketAddr> {
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<SocketAddr> {
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<SocketAddr> {
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<SocketAddr> {
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<u16>,
port_index: HashMap<Ipv4Addr, usize>,
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<u16> = (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 {
+30 -14
View File
@@ -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);
});
+6 -5
View File
@@ -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()?;
+27 -32
View File
@@ -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<AtomicCell<Instant>>,
pub head_reserve: usize,
}
impl ChannelDataHandler {
@@ -78,6 +79,7 @@ impl ChannelDataHandler {
rsa_cipher: Option<RsaCipher>,
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);
+46 -34
View File
@@ -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<Ipv4Addr> {
pub fn local_ipv4_() -> io::Result<Ipv4Addr> {
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<Ipv4Addr> {
IpAddr::V6(_) => Ok(Ipv4Addr::UNSPECIFIED),
}
}
pub fn local_ipv4() -> Option<Ipv4Addr> {
match local_ipv4_() {
Ok(ipv4) => Some(ipv4),
Err(e) => {
log::warn!("获取ipv4失败:{:?}", e);
None
}
}
}
pub fn local_ipv6() -> io::Result<Ipv6Addr> {
pub fn local_ipv6_() -> io::Result<Ipv6Addr> {
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<Ipv6Addr> {
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<Ipv6Addr> {
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<String>,
public_ip: Ipv4Addr,
public_port: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
local_ipv4: Option<Ipv4Addr>,
ipv6: Option<Ipv6Addr>,
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<Ipv4Addr>,
ipv6: Option<Ipv6Addr>,
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<String>,
public_ip: Ipv4Addr,
public_port: u16,
local_ipv4_addr: SocketAddrV4,
ipv6_addr: SocketAddrV6,
local_ipv4: Option<Ipv4Addr>,
ipv6: Option<Ipv6Addr>,
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,
)
}