diff --git a/vnt/src/handle/callback.rs b/vnt/src/handle/callback.rs new file mode 100644 index 0000000..0cc353d --- /dev/null +++ b/vnt/src/handle/callback.rs @@ -0,0 +1,201 @@ +#[cfg(feature = "server_encrypt")] +use rsa::RsaPublicKey; +use std::fmt::{Display, Formatter}; +use std::io; +use std::net::{Ipv4Addr, SocketAddr}; + +#[derive(Debug)] +pub struct DeviceInfo { + pub name: String, + pub version: String, +} + +impl Display for DeviceInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.write_str(&format!("name={} ,version={}", self.name, self.version)) + } +} + +impl DeviceInfo { + pub fn new(name: String, version: String) -> Self { + return Self { name, version }; + } +} + +#[derive(Debug)] +pub struct ConnectInfo { + // 第几次连接,从1开始 + pub count: usize, + // 服务端地址 + pub address: SocketAddr, +} + +impl Display for ConnectInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.write_str(&format!("count={} ,address={}", self.count, self.address)) + } +} + +impl ConnectInfo { + pub fn new(count: usize, address: SocketAddr) -> Self { + Self { count, address } + } +} + +#[derive(Debug)] +pub struct HandshakeInfo { + //服务端公钥 + #[cfg(feature = "server_encrypt")] + pub public_key: Option, + //服务端指纹 + #[cfg(feature = "server_encrypt")] + pub finger: Option, + //服务端版本 + pub version: String, +} + +impl Display for HandshakeInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + #[cfg(feature = "server_encrypt")] + return match &self.finger { + None => f.write_str(&format!("no_secret server version={}", self.version)), + Some(finger) => f.write_str(&format!( + "finger={} ,server version={}", + finger, self.version + )), + }; + #[cfg(not(feature = "server_encrypt"))] + f.write_str(&format!("server version={}", self.version)) + } +} +#[cfg(feature = "server_encrypt")] +impl HandshakeInfo { + pub fn new(public_key: RsaPublicKey, finger: String, version: String) -> Self { + Self { + public_key: Some(public_key), + finger: Some(finger), + version, + } + } + pub fn new_no_secret(version: String) -> Self { + Self { + public_key: None, + finger: None, + version, + } + } +} +#[cfg(not(feature = "server_encrypt"))] +impl HandshakeInfo { + pub fn new_no_secret(version: String) -> Self { + Self { version } + } +} + +#[derive(Debug)] +pub struct RegisterInfo { + //本机虚拟IP + pub virtual_ip: Ipv4Addr, + //子网掩码 + pub virtual_netmask: Ipv4Addr, + //虚拟网关 + pub virtual_gateway: Ipv4Addr, +} + +impl Display for RegisterInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.write_str(&format!( + "ip={} ,netmask={} ,gateway={}", + self.virtual_ip, self.virtual_netmask, self.virtual_gateway, + )) + } +} + +impl RegisterInfo { + pub fn new(virtual_ip: Ipv4Addr, virtual_netmask: Ipv4Addr, virtual_gateway: Ipv4Addr) -> Self { + Self { + virtual_ip, + virtual_netmask, + virtual_gateway, + } + } +} + +#[derive(Debug)] +pub struct ErrorInfo { + pub code: ErrorType, + pub msg: Option, + pub source: Option, +} + +impl Display for ErrorInfo { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.write_str(&format!("ErrorType={:?} ", self.code))?; + if let Some(msg) = &self.msg { + f.write_str(&format!(",msg={:?} ", msg))?; + } + if let Some(source) = &self.source { + f.write_str(&format!(",source={:?} ", source))?; + } + Ok(()) + } +} + +impl ErrorInfo { + pub fn new(code: ErrorType) -> Self { + Self { + code, + msg: None, + source: None, + } + } + pub fn new_msg(code: ErrorType, msg: String) -> Self { + Self { + code, + msg: Some(msg), + source: None, + } + } +} + +#[derive(Copy, Clone, Debug, Eq, PartialEq)] +pub enum ErrorType { + TokenError, + Disconnect, + AddressExhausted, + IpAlreadyExists, + InvalidIp, + Unknown, +} + +impl Into for ErrorType { + fn into(self) -> u8 { + match self { + ErrorType::TokenError => 1, + ErrorType::Disconnect => 2, + ErrorType::AddressExhausted => 3, + ErrorType::IpAlreadyExists => 4, + ErrorType::InvalidIp => 5, + ErrorType::Unknown => 6, + } + } +} + +pub trait VntCallback: Clone + Send + Sync + 'static { + /// 创建网卡的信息 + fn create_tun(&self, _info: DeviceInfo) {} + /// 连接 + fn connect(&self, _info: ConnectInfo) {} + /// 握手,返回false则拒绝握手,可在此处检查服务端信息 + fn handshake(&self, _info: HandshakeInfo) -> bool { + true + } + /// 注册,返回false则拒绝注册 + fn register(&self, _info: RegisterInfo) -> bool { + true + } + /// 异常信息 + fn error(&self, _info: ErrorInfo) {} + /// 服务停止 + fn stop(&self) {} +} diff --git a/vnt/src/handle/handshaker.rs b/vnt/src/handle/handshaker.rs new file mode 100644 index 0000000..840bf55 --- /dev/null +++ b/vnt/src/handle/handshaker.rs @@ -0,0 +1,72 @@ +use std::io; + +#[cfg(feature = "server_encrypt")] +use crate::cipher::RsaCipher; +use crate::handle::{GATEWAY_IP, SELF_IP}; +use crate::proto::message::{HandshakeRequest, SecretHandshakeRequest}; +use crate::protocol::body::RSA_ENCRYPTION_RESERVED; +use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL}; +use protobuf::Message; + +pub enum HandshakeEnum { + NotSecret, + KeyError, + Timeout, + ServerError(String), + Other(String), +} + +/// 第一次握手数据 +pub fn handshake_request_packet(secret: bool) -> io::Result>> { + let mut request = HandshakeRequest::new(); + request.secret = secret; + request.version = crate::VNT_VERSION.to_string(); + let bytes = request.write_to_bytes().map_err(|e| { + io::Error::new( + io::ErrorKind::Other, + format!("handshake_request_packet {:?}", e), + ) + })?; + let buf = vec![0u8; 12 + bytes.len()]; + let mut net_packet = NetPacket::new(buf)?; + net_packet.set_version(Version::V1); + net_packet.set_gateway_flag(true); + net_packet.set_destination(GATEWAY_IP); + net_packet.set_source(SELF_IP); + net_packet.set_protocol(Protocol::Service); + net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into()); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_payload(&bytes)?; + Ok(net_packet) +} + +/// 第二次加密握手 +#[cfg(feature = "server_encrypt")] +pub fn secret_handshake_request_packet( + rsa_cipher: &RsaCipher, + token: String, + key: &[u8], +) -> io::Result>> { + let mut request = SecretHandshakeRequest::new(); + request.token = token; + request.key = key.to_vec(); + let bytes = request.write_to_bytes().map_err(|e| { + io::Error::new( + io::ErrorKind::Other, + format!("secret_handshake_request_packet {:?}", e), + ) + })?; + let mut net_packet = NetPacket::new0( + 12 + bytes.len(), + vec![0u8; 12 + bytes.len() + RSA_ENCRYPTION_RESERVED], + )?; + net_packet.set_version(Version::V1); + net_packet.set_gateway_flag(true); + net_packet.set_destination(GATEWAY_IP); + net_packet.set_source(SELF_IP); + net_packet.set_protocol(Protocol::Service); + net_packet.set_transport_protocol(service_packet::Protocol::SecretHandshakeRequest.into()); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_payload(&bytes)?; + Ok(rsa_cipher.encrypt(&mut net_packet)?) +} diff --git a/vnt/src/handle/maintain/addr_request.rs b/vnt/src/handle/maintain/addr_request.rs new file mode 100644 index 0000000..3a11cad --- /dev/null +++ b/vnt/src/handle/maintain/addr_request.rs @@ -0,0 +1,73 @@ +use std::net::ToSocketAddrs; +use std::sync::Arc; +use std::time::Duration; + +use crossbeam_utils::atomic::AtomicCell; + +use crate::channel::context::Context; +use crate::cipher::Cipher; +use crate::handle::{BaseConfigInfo, CurrentDeviceInfo}; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL}; +use crate::util::Scheduler; + +pub fn addr_request( + scheduler: &Scheduler, + context: Context, + current_device_info: Arc>, + server_cipher: Cipher, + config: BaseConfigInfo, +) { + addr_request0(&context, ¤t_device_info, &server_cipher, &config); + // 9秒发送一次 + let rs = scheduler.timeout(Duration::from_secs(9), |s| { + addr_request(s, context, current_device_info, server_cipher, config) + }); + if !rs { + log::info!("定时任务停止"); + } +} + +pub fn addr_request0( + context: &Context, + current_device: &AtomicCell, + server_cipher: &Cipher, + config: &BaseConfigInfo, +) { + let mut current_dev = current_device.load(); + // 探测服务端地址变化 + if let Ok(mut addr) = config.server_addr.to_socket_addrs() { + if let Some(addr) = addr.next() { + if addr != current_dev.connect_server { + let mut tmp = current_dev.clone(); + tmp.connect_server = addr; + let rs = current_device.compare_exchange(current_dev, tmp); + current_dev.connect_server = addr; + log::info!( + "服务端地址变化,旧地址:{},新地址:{},替换结果:{}", + current_dev.connect_server, + addr, + rs.is_ok() + ); + } + } + } + if current_dev.connect_server.is_ipv4() { + // 如果连接的是ipv4服务,则探测公网端口 + let gateway_ip = current_dev.virtual_gateway; + let src_ip = current_dev.virtual_ip; + let mut packet = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED]).unwrap(); + packet.set_version(Version::V1); + packet.set_gateway_flag(true); + packet.set_protocol(Protocol::Control); + packet.set_transport_protocol(control_packet::Protocol::AddrRequest.into()); + packet.first_set_ttl(MAX_TTL); + packet.set_source(src_ip); + packet.set_destination(gateway_ip); + if let Err(e) = server_cipher.encrypt_ipv4(&mut packet) { + log::warn!("AddrRequest err={:?}", e) + } else { + context.try_send_all_main(packet.buffer(), current_dev.connect_server); + } + } +} diff --git a/vnt/src/handle/maintain/heartbeat.rs b/vnt/src/handle/maintain/heartbeat.rs new file mode 100644 index 0000000..968dab1 --- /dev/null +++ b/vnt/src/handle/maintain/heartbeat.rs @@ -0,0 +1,239 @@ +use std::io; +use std::net::Ipv4Addr; +use std::sync::Arc; +use std::time::Duration; + +use crossbeam_utils::atomic::AtomicCell; +use parking_lot::Mutex; +use rand::prelude::SliceRandom; + +use crate::channel::context::Context; +use crate::cipher::Cipher; +use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::control_packet::PingPacket; +use crate::protocol::{control_packet, NetPacket, Protocol, Version, MAX_TTL}; +use crate::util::Scheduler; + +/// 定时发送心跳包 +pub fn heartbeat( + scheduler: &Scheduler, + context: Context, + current_device_info: Arc>, + device_list: Arc)>>, + client_cipher: Cipher, + server_cipher: Cipher, +) { + heartbeat0( + &context, + ¤t_device_info.load(), + &device_list, + &client_cipher, + &server_cipher, + ); + // 心跳包 3秒发送一次 + let rs = scheduler.timeout(Duration::from_secs(3), |s| { + heartbeat( + s, + context, + current_device_info, + device_list, + client_cipher, + server_cipher, + ) + }); + if !rs { + log::info!("定时任务停止"); + } +} + +fn heartbeat0( + context: &Context, + current_device: &CurrentDeviceInfo, + device_list: &Mutex<(u16, Vec)>, + client_cipher: &Cipher, + server_cipher: &Cipher, +) { + let gateway_ip = current_device.virtual_gateway; + let src_ip = current_device.virtual_ip; + // 可能服务器ip发生变化,导致发送失败 + let mut is_send_gateway = false; + match heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip) { + Ok(net_packet) => { + if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server) + { + log::warn!("heartbeat err={:?}", e) + } else { + is_send_gateway = true + } + } + Err(e) => { + log::error!("heartbeat_packet err={:?}", e); + } + } + + for (dest_ip, routes) in context.route_table.route_table() { + let net_packet = if current_device.is_gateway(&dest_ip) { + if is_send_gateway { + continue; + } + heartbeat_packet_server(device_list, server_cipher, src_ip, gateway_ip) + } else { + heartbeat_packet_client(client_cipher, src_ip, dest_ip) + }; + let net_packet = match net_packet { + Ok(net_packet) => net_packet, + Err(e) => { + log::error!("heartbeat_packet err={:?}", e); + continue; + } + }; + for route in routes { + if let Err(e) = context.send_by_key(net_packet.buffer(), route.route_key()) { + log::warn!("heartbeat err={:?}", e) + } + } + } + let peer_list = { device_list.lock().1.clone() }; + for peer in &peer_list { + if !peer.status.is_online() { + continue; + } + if current_device.is_gateway(&peer.virtual_ip) { + continue; + } + if context.route_table.route_one(&peer.virtual_ip).is_none() { + //路由为空,则向服务端地址发送 + let net_packet = match heartbeat_packet_client(client_cipher, src_ip, peer.virtual_ip) { + Ok(net_packet) => net_packet, + Err(e) => { + log::error!("heartbeat_packet err={:?}", e); + continue; + } + }; + if let Err(e) = context.send_default(net_packet.buffer(), current_device.connect_server) + { + log::error!("heartbeat_packet send_default err={:?}", e); + } + } + } +} + +/// 客户端中继路径探测,延迟启动 +pub fn client_relay( + scheduler: &Scheduler, + context: Context, + current_device: Arc>, + device_list: Arc)>>, + client_cipher: Cipher, +) { + let rs = scheduler.timeout(Duration::from_secs(30), move |s| { + client_relay_(s, context, current_device, device_list, client_cipher) + }); + if !rs { + log::info!("定时任务停止"); + } +} + +/// 客户端中继路径探测,每30秒探测一次 +fn client_relay_( + scheduler: &Scheduler, + context: Context, + current_device: Arc>, + device_list: Arc)>>, + client_cipher: Cipher, +) { + if let Err(e) = client_relay0( + &context, + ¤t_device.load(), + &device_list, + &client_cipher, + ) { + log::error!("{:?}", e); + } + let rs = scheduler.timeout(Duration::from_secs(30), move |s| { + client_relay_(s, context, current_device, device_list, client_cipher) + }); + if !rs { + log::info!("定时任务停止"); + } +} + +fn client_relay0( + context: &Context, + current_device: &CurrentDeviceInfo, + device_list: &Mutex<(u16, Vec)>, + client_cipher: &Cipher, +) -> io::Result<()> { + let peer_list = { device_list.lock().1.clone() }; + let mut routes = context.route_table.route_table_p2p(); + for peer in &peer_list { + if peer.virtual_ip == current_device.virtual_ip { + continue; + } + if let Some(route) = context.route_table.route_one(&peer.virtual_ip) { + if route.is_p2p() && !context.first_latency() { + continue; + } + } + let client_packet = + heartbeat_packet_client(client_cipher, current_device.virtual_ip, peer.virtual_ip)?; + + //随机发送到其他地址,看有没有客户端符合转发条件 + routes.shuffle(&mut rand::thread_rng()); + + for (index, (ip, route)) in routes.iter().enumerate() { + if current_device.is_gateway(ip) { + continue; + } + if let Err(e) = context.send_by_key(client_packet.buffer(), route.route_key()) { + log::error!("{:?}", e); + } + if index >= 2 { + break; + } + } + } + Ok(()) +} + +/// 构建心跳包 +fn heartbeat_packet( + src: Ipv4Addr, + dest: Ipv4Addr, +) -> io::Result> { + let mut net_packet = NetPacket::new_encrypt([0u8; 12 + 4 + ENCRYPTION_RESERVED])?; + net_packet.set_version(Version::V1); + net_packet.set_protocol(Protocol::Control); + net_packet.set_transport_protocol(control_packet::Protocol::Ping.into()); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_source(src); + net_packet.set_destination(dest); + let mut ping = PingPacket::new(net_packet.payload_mut())?; + ping.set_time(crate::handle::now_time() as u16); + Ok(net_packet) +} + +fn heartbeat_packet_client( + client_cipher: &Cipher, + src: Ipv4Addr, + dest: Ipv4Addr, +) -> io::Result> { + let mut net_packet = heartbeat_packet(src, dest)?; + client_cipher.encrypt_ipv4(&mut net_packet)?; + Ok(net_packet) +} + +fn heartbeat_packet_server( + device_list: &Mutex<(u16, Vec)>, + server_cipher: &Cipher, + src: Ipv4Addr, + dest: Ipv4Addr, +) -> io::Result> { + let mut net_packet = heartbeat_packet(src, dest)?; + let mut ping = PingPacket::new(net_packet.payload_mut())?; + ping.set_epoch(device_list.lock().0); + net_packet.set_gateway_flag(true); + server_cipher.encrypt_ipv4(&mut net_packet)?; + Ok(net_packet) +} diff --git a/vnt/src/handle/maintain/idle.rs b/vnt/src/handle/maintain/idle.rs new file mode 100644 index 0000000..a1bc82c --- /dev/null +++ b/vnt/src/handle/maintain/idle.rs @@ -0,0 +1,133 @@ +use crate::channel::context::Context; +use crate::channel::idle::{Idle, IdleType}; +use crate::channel::sender::AcceptSocketSender; +use crate::handle::callback::{ConnectInfo, ErrorType}; +use crate::handle::{handshaker, BaseConfigInfo, CurrentDeviceInfo}; +use crate::util::Scheduler; +use crate::{ErrorInfo, VntCallback}; +use crossbeam_utils::atomic::AtomicCell; +use mio::net::TcpStream; +use std::io; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +pub fn idle_route( + scheduler: &Scheduler, + idle: Idle, + context: Context, + current_device_info: Arc>, + call: Call, +) { + let delay = idle_route0(&idle, &context, ¤t_device_info, &call); + let rs = scheduler.timeout(delay, move |s| { + idle_route(s, idle, context, current_device_info, call) + }); + if !rs { + log::info!("定时任务停止"); + } +} +pub fn idle_gateway( + scheduler: &Scheduler, + context: Context, + current_device_info: Arc>, + config: BaseConfigInfo, + tcp_socket_sender: AcceptSocketSender<(TcpStream, SocketAddr, Option>)>, + call: Call, + mut connect_count: usize, +) { + idle_gateway0( + &context, + ¤t_device_info, + &config, + &tcp_socket_sender, + &call, + &mut connect_count, + ); + let rs = scheduler.timeout(Duration::from_secs(5), move |s| { + idle_gateway( + s, + context, + current_device_info, + config, + tcp_socket_sender, + call, + connect_count, + ) + }); + if !rs { + log::info!("定时任务停止"); + } +} +fn idle_gateway0( + context: &Context, + current_device: &AtomicCell, + config: &BaseConfigInfo, + tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option>)>, + call: &Call, + connect_count: &mut usize, +) { + let cur = current_device.load(); + if let Err(e) = + check_gateway_channel(context, cur, config, tcp_socket_sender, call, connect_count) + { + log::warn!("{:?}", e); + } +} +fn idle_route0( + idle: &Idle, + context: &Context, + current_device: &AtomicCell, + call: &Call, +) -> Duration { + let cur = current_device.load(); + match idle.next_idle() { + IdleType::Timeout(ip, route) => { + context.route_table.remove_route(&ip, route); + if cur.is_gateway(&ip) { + //网关路由过期,则需要改变状态 + let _ = context.change_status(current_device); + call.error(ErrorInfo::new(ErrorType::Disconnect)); + } + Duration::from_millis(100) + } + IdleType::Sleep(duration) => duration, + IdleType::None => Duration::from_millis(3000), + } +} + +fn check_gateway_channel( + context: &Context, + current_device: CurrentDeviceInfo, + config: &BaseConfigInfo, + tcp_socket_sender: &AcceptSocketSender<(TcpStream, SocketAddr, Option>)>, + call: &Call, + count: &mut usize, +) -> io::Result<()> { + let gateway_route = context + .route_table + .route_one(¤t_device.virtual_gateway); + if gateway_route.is_none() { + *count += 1; + //需要重连 + call.connect(ConnectInfo::new(*count, current_device.connect_server)); + let request_packet = handshaker::handshake_request_packet(config.client_secret)?; + if let Err(e) = context.send_default(request_packet.buffer(), current_device.connect_server) + { + log::warn!("{:?}", e); + if context.is_main_tcp() { + //tcp需要重连 + let tcp_stream = std::net::TcpStream::connect(current_device.connect_server)?; + tcp_stream.set_nonblocking(true)?; + if let Err(e) = tcp_socket_sender.try_add_socket(( + TcpStream::from_std(tcp_stream), + current_device.connect_server, + Some(request_packet.into_buffer()), + )) { + log::warn!("{:?}", e) + } + } + } + } + Ok(()) +} diff --git a/vnt/src/handle/maintain/mod.rs b/vnt/src/handle/maintain/mod.rs new file mode 100644 index 0000000..b685fe4 --- /dev/null +++ b/vnt/src/handle/maintain/mod.rs @@ -0,0 +1,16 @@ +mod heartbeat; +pub use heartbeat::client_relay; +pub use heartbeat::heartbeat; + +mod re_nat_type; +pub use re_nat_type::retrieve_nat_type; + +mod addr_request; +pub use addr_request::addr_request; + +mod punch; +pub use punch::punch; + +mod idle; +pub use idle::idle_gateway; +pub use idle::idle_route; diff --git a/vnt/src/handle/maintain/punch.rs b/vnt/src/handle/maintain/punch.rs new file mode 100644 index 0000000..23ee1d7 --- /dev/null +++ b/vnt/src/handle/maintain/punch.rs @@ -0,0 +1,185 @@ +use std::net::Ipv4Addr; +use std::sync::mpsc::Receiver; +use std::sync::Arc; +use std::time::Duration; +use std::{io, thread}; + +use crossbeam_utils::atomic::AtomicCell; +use parking_lot::Mutex; +use protobuf::Message; +use rand::prelude::SliceRandom; + +use crate::channel::context::Context; +use crate::channel::punch::{NatInfo, Punch}; +use crate::cipher::Cipher; +use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; +use crate::nat::NatTest; +use crate::proto::message::{PunchInfo, PunchNatType}; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::{control_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL}; +use crate::util::Scheduler; + +pub fn punch( + scheduler: &Scheduler, + context: Context, + nat_test: NatTest, + device_list: Arc)>>, + current_device: Arc>, + client_cipher: Cipher, + receiver: Receiver<(Ipv4Addr, NatInfo)>, + punch: Punch, +) { + punch_request( + scheduler, + context, + nat_test, + device_list, + current_device.clone(), + client_cipher.clone(), + 0, + ); + thread::spawn(move || { + punch_start(receiver, punch, current_device, client_cipher); + }); +} + +/// 接收打洞消息,配合对端打洞 +fn punch_start( + receiver: Receiver<(Ipv4Addr, NatInfo)>, + mut punch: Punch, + current_device: Arc>, + client_cipher: Cipher, +) { + while let Ok((peer_ip, nat_info)) = receiver.recv() { + let mut packet = NetPacket::new_encrypt([0u8; 12 + ENCRYPTION_RESERVED]).unwrap(); + 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.load().virtual_ip()); + packet.set_destination(peer_ip); + log::info!("发起打洞,目标:{:?},{:?}", peer_ip, nat_info); + if let Err(e) = client_cipher.encrypt_ipv4(&mut packet) { + log::error!("{:?}", e); + continue; + } + if let Err(e) = punch.punch(packet.buffer(), peer_ip, nat_info) { + log::warn!("{:?}", e) + } + } +} + +/// 定时发起打洞请求 +fn punch_request( + scheduler: &Scheduler, + context: Context, + nat_test: NatTest, + device_list: Arc)>>, + current_device: Arc>, + client_cipher: Cipher, + count: usize, +) { + if let Err(e) = punch0( + &context, + &nat_test, + &device_list, + ¤t_device, + &client_cipher, + ) { + log::warn!("{:?}", e) + } + let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29]; + let secs = Duration::from_secs(sleep_time[count % sleep_time.len()]); + let rs = scheduler.timeout(secs, move |s| { + punch_request( + s, + context, + nat_test, + device_list, + current_device, + client_cipher, + count + 1, + ); + }); + if !rs { + log::info!("定时任务停止"); + } +} + +/// 随机对需要打洞的客户端发起打洞请求 +fn punch0( + context: &Context, + nat_test: &NatTest, + device_list: &Arc)>>, + current_device: &Arc>, + client_cipher: &Cipher, +) -> io::Result<()> { + let current_device = current_device.load(); + let nat_info = nat_test.nat_info(); + let mut list = device_list.lock().clone().1; + list.shuffle(&mut rand::thread_rng()); + let mut count = 0; + for info in list { + if !info.status.is_online() { + continue; + } + if info.virtual_ip <= current_device.virtual_ip { + continue; + } + if !context.route_table.need_punch(&info.virtual_ip) { + continue; + } + count += 1; + if count > 2 { + break; + } + let packet = punch_packet( + client_cipher, + current_device.virtual_ip(), + &nat_info, + info.virtual_ip, + )?; + context.send_default(packet.buffer(), current_device.connect_server)?; + } + Ok(()) +} + +fn punch_packet( + client_cipher: &Cipher, + virtual_ip: Ipv4Addr, + nat_info: &NatInfo, + dest: Ipv4Addr, +) -> io::Result>> { + let mut punch_reply = PunchInfo::new(); + punch_reply.reply = false; + punch_reply.public_ip_list = nat_info + .public_ips + .iter() + .map(|ip| u32::from_be_bytes(ip.octets())) + .collect(); + punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32); + punch_reply.public_ports = nat_info.public_ports.iter().map(|e| *e as u32).collect(); + punch_reply.public_port_range = nat_info.public_port_range as u32; + punch_reply.local_ip = u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED)); + punch_reply.local_port = nat_info.udp_ports[0] as u32; + punch_reply.tcp_port = nat_info.tcp_port as u32; + punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect(); + if let Some(ipv6) = nat_info.ipv6 { + punch_reply.ipv6_port = nat_info.udp_ports[0] 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() + .map_err(|e| io::Error::new(io::ErrorKind::Other, format!("punch_packet {:?}", e)))?; + let mut net_packet = NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; + net_packet.set_version(Version::V1); + net_packet.set_protocol(Protocol::OtherTurn); + net_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into()); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_source(virtual_ip); + net_packet.set_destination(dest); + net_packet.set_payload(&bytes)?; + client_cipher.encrypt_ipv4(&mut net_packet)?; + Ok(net_packet) +} diff --git a/vnt/src/handle/maintain/re_nat_type.rs b/vnt/src/handle/maintain/re_nat_type.rs new file mode 100644 index 0000000..7645626 --- /dev/null +++ b/vnt/src/handle/maintain/re_nat_type.rs @@ -0,0 +1,45 @@ +use std::thread; +use std::time::Duration; + +use crate::channel::context::Context; +use crate::channel::sender::AcceptSocketSender; +use crate::nat; +use crate::nat::NatTest; +use crate::util::Scheduler; + +/// 10分钟探测一次nat +pub fn retrieve_nat_type( + scheduler: &Scheduler, + context: Context, + nat_test: NatTest, + udp_socket_sender: AcceptSocketSender>>, +) { + retrieve_nat_type0(context.clone(), nat_test.clone(), udp_socket_sender.clone()); + scheduler.timeout(Duration::from_secs(60 * 10), move |s| { + retrieve_nat_type(s, context, nat_test, udp_socket_sender) + }); +} + +fn retrieve_nat_type0( + context: Context, + nat_test: NatTest, + udp_socket_sender: AcceptSocketSender>>, +) { + thread::spawn(move || { + if nat_test.can_update() { + let nat_info = nat_test.nat_info(); + let local_ipv4 = nat::local_ipv4(); + let local_ipv6 = nat::local_ipv6(); + let nat_info = nat_test.re_test( + nat_info.public_ports, + local_ipv4, + local_ipv6, + nat_info.udp_ports, + nat_info.tcp_port, + ); + if let Err(e) = context.switch(nat_info.nat_type, &udp_socket_sender) { + log::warn!("{:?}", e); + } + } + }); +} diff --git a/vnt/src/handle/mod.rs b/vnt/src/handle/mod.rs index 5c39f43..28db419 100644 --- a/vnt/src/handle/mod.rs +++ b/vnt/src/handle/mod.rs @@ -1,12 +1,15 @@ use std::net::{Ipv4Addr, SocketAddr}; -pub mod handshake_handler; -pub mod heartbeat_handler; -pub mod punch_handler; -pub mod recv_handler; -pub mod registration_handler; +pub mod callback; +pub mod handshaker; +pub mod maintain; +pub mod recv_data; +pub mod registrar; pub mod tun_tap; +const SELF_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 2); +const GATEWAY_IP: Ipv4Addr = Ipv4Addr::new(0, 0, 0, 1); + pub fn now_time() -> u64 { let now = std::time::SystemTime::now(); if let Ok(timestamp) = now.duration_since(std::time::UNIX_EPOCH) { @@ -41,12 +44,48 @@ impl PeerDeviceInfo { } } +#[derive(Clone, Debug)] +pub struct BaseConfigInfo { + pub name: String, + pub token: String, + pub ip: Option, + pub client_secret: bool, + pub device_id: String, + pub server_addr: String, +} + +impl BaseConfigInfo { + pub fn new( + name: String, + token: String, + ip: Option, + client_secret: bool, + device_id: String, + server_addr: String, + ) -> Self { + Self { + name, + token, + ip, + client_secret, + device_id, + server_addr, + } + } +} + #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd)] pub enum PeerDeviceStatus { Online, Offline, } +impl PeerDeviceStatus { + pub fn is_online(&self) -> bool { + self == &PeerDeviceStatus::Online + } +} + impl Into for PeerDeviceStatus { fn into(self) -> u8 { match self { @@ -73,27 +112,32 @@ pub enum ConnectStatus { #[derive(Copy, Clone, Debug, Eq, PartialEq)] pub struct CurrentDeviceInfo { - virtual_ip: Ipv4Addr, - pub virtual_gateway: Ipv4Addr, + //本机虚拟IP + pub virtual_ip: Ipv4Addr, + //子网掩码 pub virtual_netmask: Ipv4Addr, + //虚拟网关 + pub virtual_gateway: Ipv4Addr, //网络地址 pub virtual_network: Ipv4Addr, //直接广播地址 - pub broadcast_address: Ipv4Addr, + pub broadcast_ip: Ipv4Addr, //链接的服务器地址 pub connect_server: SocketAddr, + //连接状态 + pub status: ConnectStatus, } impl CurrentDeviceInfo { pub fn new( virtual_ip: Ipv4Addr, - virtual_gateway: Ipv4Addr, virtual_netmask: Ipv4Addr, + virtual_gateway: Ipv4Addr, connect_server: SocketAddr, ) -> Self { - let broadcast_address = (!u32::from_be_bytes(virtual_netmask.octets())) + let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets())) | u32::from_be_bytes(virtual_gateway.octets()); - let broadcast_address = Ipv4Addr::from(broadcast_address); + let broadcast_ip = Ipv4Addr::from(broadcast_ip); let virtual_network = u32::from_be_bytes(virtual_netmask.octets()) & u32::from_be_bytes(virtual_gateway.octets()); let virtual_network = Ipv4Addr::from(virtual_network); @@ -102,10 +146,40 @@ impl CurrentDeviceInfo { virtual_netmask, virtual_gateway, virtual_network, - broadcast_address, + broadcast_ip, connect_server, + status: ConnectStatus::Connecting, } } + pub fn new0(connect_server: SocketAddr) -> Self { + Self { + virtual_ip: Ipv4Addr::UNSPECIFIED, + virtual_gateway: Ipv4Addr::UNSPECIFIED, + virtual_netmask: Ipv4Addr::UNSPECIFIED, + virtual_network: Ipv4Addr::UNSPECIFIED, + broadcast_ip: Ipv4Addr::UNSPECIFIED, + connect_server, + status: ConnectStatus::Connecting, + } + } + pub fn update( + &mut self, + virtual_ip: Ipv4Addr, + virtual_netmask: Ipv4Addr, + virtual_gateway: Ipv4Addr, + ) { + let broadcast_ip = (!u32::from_be_bytes(virtual_netmask.octets())) + | u32::from_be_bytes(virtual_gateway.octets()); + let broadcast_ip = Ipv4Addr::from(broadcast_ip); + let virtual_network = u32::from_be_bytes(virtual_netmask.octets()) + & u32::from_be_bytes(virtual_gateway.octets()); + let virtual_network = Ipv4Addr::from(virtual_network); + self.virtual_ip = virtual_ip; + self.virtual_netmask = virtual_netmask; + self.virtual_gateway = virtual_gateway; + self.broadcast_ip = broadcast_ip; + self.virtual_network = virtual_network; + } #[inline] pub fn virtual_ip(&self) -> Ipv4Addr { self.virtual_ip @@ -114,4 +188,7 @@ impl CurrentDeviceInfo { pub fn virtual_gateway(&self) -> Ipv4Addr { self.virtual_gateway } + pub fn is_gateway(&self, ip: &Ipv4Addr) -> bool { + &self.virtual_gateway == ip || ip == &GATEWAY_IP + } } diff --git a/vnt/src/handle/recv_data/client.rs b/vnt/src/handle/recv_data/client.rs new file mode 100644 index 0000000..5470242 --- /dev/null +++ b/vnt/src/handle/recv_data/client.rs @@ -0,0 +1,338 @@ +use std::collections::HashMap; +use std::io; +use std::net::{Ipv4Addr, Ipv6Addr}; +use std::sync::mpsc::SyncSender; +use std::sync::Arc; + +use parking_lot::RwLock; +use protobuf::Message; + +use packet::icmp::{icmp, Kind}; +use packet::ip::ipv4; +use packet::ip::ipv4::packet::IpV4Packet; +use tun::device::IFace; +use tun::Device; + +use crate::channel::context::Context; +use crate::channel::punch::NatInfo; +use crate::channel::{Route, RouteKey}; +use crate::cipher::Cipher; +use crate::external_route::AllowExternalRoute; +use crate::handle::recv_data::PacketHandler; +use crate::handle::CurrentDeviceInfo; +#[cfg(feature = "ip_proxy")] +use crate::ip_proxy::{IpProxyMap, ProxyHandler}; +use crate::nat::NatTest; +use crate::proto::message::{PunchInfo, PunchNatType}; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::control_packet::ControlPacket; +use crate::protocol::{ + control_packet, ip_turn_packet, other_turn_packet, NetPacket, Protocol, Version, MAX_TTL, +}; + +/// 处理来源于客户端的包 +#[derive(Clone)] +pub struct ClientPacketHandler { + device: Arc, + client_cipher: Cipher, + relay: bool, + punch_sender: SyncSender<(Ipv4Addr, NatInfo)>, + peer_nat_info_map: Arc>>, + nat_test: NatTest, + route: AllowExternalRoute, + #[cfg(feature = "ip_proxy")] + ip_proxy_map: Option, +} + +impl ClientPacketHandler { + pub fn new( + device: Arc, + client_cipher: Cipher, + relay: bool, + punch_sender: SyncSender<(Ipv4Addr, NatInfo)>, + peer_nat_info_map: Arc>>, + nat_test: NatTest, + route: AllowExternalRoute, + #[cfg(feature = "ip_proxy")] ip_proxy_map: Option, + ) -> Self { + Self { + device, + client_cipher, + relay, + punch_sender, + peer_nat_info_map, + nat_test, + route, + #[cfg(feature = "ip_proxy")] + ip_proxy_map, + } + } +} + +impl PacketHandler for ClientPacketHandler { + fn handle( + &self, + mut net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + context: &Context, + current_device: &CurrentDeviceInfo, + ) -> io::Result<()> { + self.client_cipher.decrypt_ipv4(&mut net_packet)?; + match net_packet.protocol() { + Protocol::Service => {} + Protocol::Error => {} + Protocol::Control => { + self.control(context, current_device, net_packet, route_key)?; + } + Protocol::IpTurn => { + self.ip_turn(net_packet, context, current_device, route_key)?; + } + Protocol::OtherTurn => { + self.other_turn(context, current_device, net_packet, route_key)?; + } + Protocol::Unknown(_) => {} + } + Ok(()) + } +} + +impl ClientPacketHandler { + fn ip_turn( + &self, + mut net_packet: NetPacket<&mut [u8]>, + context: &Context, + current_device: &CurrentDeviceInfo, + route_key: RouteKey, + ) -> io::Result<()> { + let destination = net_packet.destination(); + let source = net_packet.source(); + match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) { + ip_turn_packet::Protocol::Ipv4 => { + let mut ipv4 = IpV4Packet::new(net_packet.payload_mut())?; + match ipv4.protocol() { + ipv4::protocol::Protocol::Icmp => { + if ipv4.destination_ip() == destination { + let mut icmp_packet = icmp::IcmpPacket::new(ipv4.payload_mut())?; + if icmp_packet.kind() == Kind::EchoRequest { + //开启ping + icmp_packet.set_kind(Kind::EchoReply); + icmp_packet.update_checksum(); + ipv4.set_source_ip(destination); + ipv4.set_destination_ip(source); + ipv4.update_checksum(); + net_packet.set_source(destination); + net_packet.set_destination(source); + //不管加不加密,和接收到的数据长度都一致 + self.client_cipher.encrypt_ipv4(&mut net_packet)?; + context.send_by_key(net_packet.buffer(), route_key)?; + return Ok(()); + } + } + } + _ => {} + } + // ip代理只关心实际目标 + let real_dest = ipv4.destination_ip(); + if real_dest != destination + && !(real_dest.is_broadcast() + || real_dest.is_multicast() + || real_dest == current_device.broadcast_ip + || real_dest.is_unspecified()) + { + if !self.route.allow(&ipv4.destination_ip()) { + //拦截不符合的目标 + return Ok(()); + } + #[cfg(feature = "ip_proxy")] + if let Some(ip_proxy_map) = &self.ip_proxy_map { + if ip_proxy_map.recv_handle(&mut ipv4, source, destination)? { + return Ok(()); + } + } + } + self.device.write(net_packet.payload())?; + } + ip_turn_packet::Protocol::Ipv4Broadcast => { + //客户端不帮忙转发广播包,所以不会出现这种类型的数据 + } + ip_turn_packet::Protocol::Unknown(_) => {} + } + Ok(()) + } + fn control( + &self, + context: &Context, + current_device: &CurrentDeviceInfo, + mut net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + ) -> io::Result<()> { + let metric = net_packet.source_ttl() - net_packet.ttl() + 1; + let source = net_packet.source(); + match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { + ControlPacket::PingPacket(_) => { + net_packet.set_transport_protocol(control_packet::Protocol::Pong.into()); + net_packet.set_source(current_device.virtual_ip); + net_packet.set_destination(source); + net_packet.first_set_ttl(MAX_TTL); + self.client_cipher.encrypt_ipv4(&mut net_packet)?; + context.send_by_key(net_packet.buffer(), route_key)?; + let route = Route::from(route_key, metric, 199); + context.route_table.add_route_if_absent(source, route); + } + ControlPacket::PongPacket(pong_packet) => { + let current_time = crate::handle::now_time() as u16; + if current_time < pong_packet.time() { + return Ok(()); + } + let rt = (current_time - pong_packet.time()) as i64; + let route = Route::from(route_key, metric, rt); + context.route_table.add_route(source, route); + } + ControlPacket::PunchRequest => { + if self.relay { + return Ok(()); + } + //回应 + net_packet.set_transport_protocol(control_packet::Protocol::PunchResponse.into()); + net_packet.set_source(current_device.virtual_ip); + net_packet.set_destination(source); + net_packet.first_set_ttl(1); + self.client_cipher.encrypt_ipv4(&mut net_packet)?; + context.send_by_key(net_packet.buffer(), route_key)?; + let route = Route::from(route_key, 1, 199); + context.route_table.add_route_if_absent(source, route); + } + ControlPacket::PunchResponse => { + if self.relay { + return Ok(()); + } + let route = Route::from(route_key, 1, 199); + context.route_table.add_route_if_absent(source, route); + } + ControlPacket::AddrRequest => match route_key.addr.ip() { + std::net::IpAddr::V4(ipv4) => { + let mut packet = NetPacket::new_encrypt([0; 12 + 6 + ENCRYPTION_RESERVED])?; + packet.set_version(Version::V1); + packet.set_protocol(Protocol::Control); + packet.set_transport_protocol(control_packet::Protocol::AddrResponse.into()); + packet.first_set_ttl(MAX_TTL); + packet.set_source(current_device.virtual_ip); + packet.set_destination(source); + let mut addr_packet = control_packet::AddrPacket::new(packet.payload_mut())?; + addr_packet.set_ipv4(ipv4); + addr_packet.set_port(route_key.addr.port()); + self.client_cipher.encrypt_ipv4(&mut packet)?; + context.send_by_key(packet.buffer(), route_key)?; + } + std::net::IpAddr::V6(_) => {} + }, + ControlPacket::AddrResponse(_) => {} + } + Ok(()) + } + fn other_turn( + &self, + context: &Context, + current_device: &CurrentDeviceInfo, + net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + ) -> io::Result<()> { + if self.relay { + return Ok(()); + } + let source = net_packet.source(); + match other_turn_packet::Protocol::from(net_packet.transport_protocol()) { + other_turn_packet::Protocol::Punch => { + let mut punch_info = + PunchInfo::parse_from_bytes(net_packet.payload()).map_err(|e| { + io::Error::new(io::ErrorKind::Other, format!("PunchInfo {:?}", e)) + })?; + let public_ips = punch_info + .public_ip_list + .iter() + .map(|v| Ipv4Addr::from(v.to_be_bytes())) + .collect(); + let local_ipv4 = Some(Ipv4Addr::from(punch_info.local_ip.to_be_bytes())); + 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(); + Some(Ipv6Addr::from(ipv6)) + } else { + None + }; + //兼容旧版本 + if punch_info.public_ports.is_empty() { + punch_info.public_ports.push(punch_info.public_port); + } + //兼容旧版本 + if punch_info.udp_ports.is_empty() { + punch_info.udp_ports.push(punch_info.local_port); + } + let peer_nat_info = NatInfo::new( + public_ips, + punch_info.public_ports.iter().map(|e| *e as u16).collect(), + punch_info.public_port_range as u16, + local_ipv4, + ipv6, + punch_info.udp_ports.iter().map(|e| *e as u16).collect(), + tcp_port, + punch_info.nat_type.enum_value_or_default().into(), + ); + { + let peer_nat_info = peer_nat_info.clone(); + self.peer_nat_info_map.write().insert(source, peer_nat_info); + } + if !punch_info.reply { + let mut punch_reply = PunchInfo::new(); + punch_reply.reply = true; + let nat_info = self.nat_test.nat_info(); + punch_reply.public_ip_list = nat_info + .public_ips + .iter() + .map(|ip| u32::from_be_bytes(ip.octets())) + .collect(); + punch_reply.public_port = nat_info.public_ports.get(0).map_or(0, |v| *v as u32); + punch_reply.public_ports = + nat_info.public_ports.iter().map(|e| *e as u32).collect(); + punch_reply.public_port_range = nat_info.public_port_range as u32; + punch_reply.tcp_port = nat_info.tcp_port as u32; + punch_reply.nat_type = + protobuf::EnumOrUnknown::new(PunchNatType::from(nat_info.nat_type)); + punch_reply.local_ip = + u32::from(nat_info.local_ipv4().unwrap_or(Ipv4Addr::UNSPECIFIED)); + punch_reply.local_port = nat_info.udp_ports[0] as u32; + punch_reply.udp_ports = nat_info.udp_ports.iter().map(|e| *e as u32).collect(); + if let Some(ipv6) = nat_info.ipv6() { + punch_reply.ipv6 = ipv6.octets().to_vec(); + punch_reply.ipv6_port = nat_info.udp_ports[0] as u32; + } + let bytes = punch_reply.write_to_bytes().map_err(|e| { + io::Error::new(io::ErrorKind::Other, format!("punch_reply {:?}", e)) + })?; + let mut punch_packet = + NetPacket::new_encrypt(vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED])?; + punch_packet.set_version(Version::V1); + punch_packet.set_protocol(Protocol::OtherTurn); + punch_packet.set_transport_protocol(other_turn_packet::Protocol::Punch.into()); + punch_packet.first_set_ttl(MAX_TTL); + punch_packet.set_source(current_device.virtual_ip()); + punch_packet.set_destination(source); + punch_packet.set_payload(&bytes)?; + if self.punch(source, peer_nat_info) { + self.client_cipher.encrypt_ipv4(&mut punch_packet)?; + context.send_by_key(punch_packet.buffer(), route_key)?; + } + } else { + self.punch(source, peer_nat_info); + } + } + other_turn_packet::Protocol::Unknown(e) => { + log::warn!("不支持的转发协议 {:?},source:{:?}", e, source); + } + } + Ok(()) + } + fn punch(&self, peer_ip: Ipv4Addr, peer_nat_info: NatInfo) -> bool { + self.punch_sender.try_send((peer_ip, peer_nat_info)).is_ok() + } +} diff --git a/vnt/src/handle/recv_data/mod.rs b/vnt/src/handle/recv_data/mod.rs new file mode 100644 index 0000000..13a9c90 --- /dev/null +++ b/vnt/src/handle/recv_data/mod.rs @@ -0,0 +1,152 @@ +use std::collections::HashMap; +use std::net::Ipv4Addr; +use std::sync::mpsc::SyncSender; +use std::sync::Arc; +use std::{io, thread}; + +use crossbeam_utils::atomic::AtomicCell; +use parking_lot::{Mutex, RwLock}; + +use tun::Device; + +use crate::channel::context::Context; +use crate::channel::handler::RecvChannelHandler; +use crate::channel::punch::NatInfo; +use crate::channel::RouteKey; +use crate::cipher::Cipher; +#[cfg(feature = "server_encrypt")] +use crate::cipher::RsaCipher; +use crate::external_route::{AllowExternalRoute, ExternalRoute}; +use crate::handle::callback::VntCallback; +use crate::handle::recv_data::client::ClientPacketHandler; +use crate::handle::recv_data::server::ServerPacketHandler; +use crate::handle::recv_data::turn::TurnPacketHandler; +use crate::handle::{BaseConfigInfo, CurrentDeviceInfo, PeerDeviceInfo, SELF_IP}; +#[cfg(feature = "ip_proxy")] +use crate::ip_proxy::IpProxyMap; +use crate::nat::NatTest; +use crate::protocol::NetPacket; +use crate::util::U64Adder; + +mod client; +mod server; +mod turn; + +#[derive(Clone)] +pub struct RecvDataHandler { + current_device: Arc>, + turn: TurnPacketHandler, + client: ClientPacketHandler, + server: ServerPacketHandler, + counter: U64Adder, +} + +impl RecvChannelHandler for RecvDataHandler { + fn handle(&mut self, buf: &mut [u8], route_key: RouteKey, context: &Context) { + if let Err(e) = self.handle0(buf, route_key, context) { + log::error!("[{}]-{:?}", thread::current().name().unwrap_or(""), e); + } + } +} + +impl RecvDataHandler { + pub fn new( + #[cfg(feature = "server_encrypt")] rsa_cipher: Arc>>, + server_cipher: Cipher, + client_cipher: Cipher, + current_device: Arc>, + device: Arc, + device_list: Arc)>>, + config_info: BaseConfigInfo, + nat_test: NatTest, + callback: Call, + relay: bool, + punch_sender: SyncSender<(Ipv4Addr, NatInfo)>, + peer_nat_info_map: Arc>>, + external_route: ExternalRoute, + route: AllowExternalRoute, + #[cfg(feature = "ip_proxy")] ip_proxy_map: Option, + counter: U64Adder, + ) -> Self { + let server = ServerPacketHandler::new( + #[cfg(feature = "server_encrypt")] + rsa_cipher, + server_cipher, + current_device.clone(), + device.clone(), + device_list, + config_info, + nat_test.clone(), + callback, + external_route, + ); + let client = ClientPacketHandler::new( + device.clone(), + client_cipher, + relay, + punch_sender, + peer_nat_info_map, + nat_test, + route, + #[cfg(feature = "ip_proxy")] + ip_proxy_map, + ); + let turn = TurnPacketHandler::new(); + Self { + current_device, + turn, + client, + server, + counter, + } + } + fn handle0( + &mut self, + buf: &mut [u8], + route_key: RouteKey, + context: &Context, + ) -> io::Result<()> { + // 统计流量 + self.counter.add(buf.len() as _); + let net_packet = NetPacket::new(buf)?; + if net_packet.ttl() == 0 || net_packet.source_ttl() < net_packet.ttl() { + return Ok(()); + } + let current_device = self.current_device.load(); + let dest = net_packet.destination(); + let source = net_packet.source(); + context.route_table.update_read_time(&source, &route_key); + if dest == current_device.virtual_ip + || dest.is_broadcast() + || dest.is_multicast() + || dest == SELF_IP + || dest.is_unspecified() + || dest == current_device.broadcast_ip + { + //发给自己的包 + if net_packet.is_gateway() { + //服务端-客户端包 + self.server + .handle(net_packet, route_key, context, ¤t_device) + } else { + //客户端-客户端包 + self.client + .handle(net_packet, route_key, context, ¤t_device) + } + } else { + //转发包 + self.turn + .handle(net_packet, route_key, context, ¤t_device) + } + } +} + +pub trait PacketHandler { + fn handle( + &self, + net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + context: &Context, + current_device: &CurrentDeviceInfo, + ) -> io::Result<()>; +} diff --git a/vnt/src/handle/recv_data/server.rs b/vnt/src/handle/recv_data/server.rs new file mode 100644 index 0000000..654d8b8 --- /dev/null +++ b/vnt/src/handle/recv_data/server.rs @@ -0,0 +1,437 @@ +use std::io; +use std::net::Ipv4Addr; +use std::ops::Sub; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use crossbeam_utils::atomic::AtomicCell; +use parking_lot::Mutex; +use protobuf::Message; + +use packet::icmp::{icmp, Kind}; +use packet::ip::ipv4; +use packet::ip::ipv4::packet::IpV4Packet; +use tun::device::IFace; +use tun::Device; + +use crate::channel::context::Context; +use crate::channel::{Route, RouteKey}; +use crate::cipher::Cipher; +#[cfg(feature = "server_encrypt")] +use crate::cipher::RsaCipher; +use crate::external_route::ExternalRoute; +use crate::handle::callback::{ErrorInfo, ErrorType, HandshakeInfo, RegisterInfo, VntCallback}; +use crate::handle::recv_data::PacketHandler; +use crate::handle::{ + handshaker, registrar, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, + GATEWAY_IP, +}; +use crate::nat::NatTest; +use crate::proto; +use crate::proto::message::{DeviceList, HandshakeResponse, RegistrationResponse}; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::control_packet::ControlPacket; +use crate::protocol::error_packet::InErrorPacket; +use crate::protocol::{ip_turn_packet, service_packet, NetPacket, Protocol, Version, MAX_TTL}; + +/// 处理来源于服务端的包 +#[derive(Clone)] +pub struct ServerPacketHandler { + #[cfg(feature = "server_encrypt")] + rsa_cipher: Arc>>, + server_cipher: Cipher, + current_device: Arc>, + device: Arc, + device_list: Arc)>>, + config_info: BaseConfigInfo, + nat_test: NatTest, + callback: Call, + time: Arc>, + route_record: Arc>>, + external_route: ExternalRoute, +} + +impl ServerPacketHandler { + pub fn new( + #[cfg(feature = "server_encrypt")] rsa_cipher: Arc>>, + server_cipher: Cipher, + current_device: Arc>, + device: Arc, + device_list: Arc)>>, + config_info: BaseConfigInfo, + nat_test: NatTest, + callback: Call, + external_route: ExternalRoute, + ) -> Self { + Self { + #[cfg(feature = "server_encrypt")] + rsa_cipher, + server_cipher, + current_device, + device, + device_list, + config_info, + nat_test, + callback, + time: Arc::new(AtomicCell::new(Instant::now().sub(Duration::from_secs(60)))), + route_record: Arc::new(Mutex::default()), + external_route, + } + } +} + +impl PacketHandler for ServerPacketHandler { + fn handle( + &self, + mut net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + context: &Context, + current_device: &CurrentDeviceInfo, + ) -> io::Result<()> { + if net_packet.protocol() == Protocol::Error + && net_packet.transport_protocol() + == crate::protocol::error_packet::Protocol::NoKey.into() + { + //开启服务端加密的情况下,只有这个回应是不加密的,是服务端通知客户端上传密钥 + #[cfg(feature = "server_encrypt")] + { + let mutex_guard = self.rsa_cipher.lock(); + if let Some(rsa_cipher) = mutex_guard.as_ref() { + let last = self.time.load(); + if last.elapsed() < Duration::from_secs(1) + || self.time.compare_exchange(last, Instant::now()).is_err() + { + //短时间不重复上传服务端密钥 + return Ok(()); + } + if let Some(key) = self.server_cipher.key() { + log::warn!("上传密钥到服务端:{:?}", route_key); + let packet = handshaker::secret_handshake_request_packet( + rsa_cipher, + self.config_info.token.clone(), + key, + )?; + context.send_by_key(packet.buffer(), route_key)?; + } + } + } + return Ok(()); + } else if net_packet.protocol() == Protocol::Service + && net_packet.transport_protocol() == service_packet::Protocol::HandshakeResponse.into() + { + let response = + HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| { + io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e)) + })?; + //如果开启了加密,则发送加密握手请求 + #[cfg(feature = "server_encrypt")] + if let Some(key) = self.server_cipher.key() { + let rsa_cipher = RsaCipher::new(&response.public_key)?; + let handshake_info = HandshakeInfo::new( + rsa_cipher.public_key()?.clone(), + rsa_cipher.finger()?, + response.version, + ); + log::warn!("加密握手请求:{:?}", handshake_info); + + if self.callback.handshake(handshake_info) { + let packet = handshaker::secret_handshake_request_packet( + &rsa_cipher, + self.config_info.token.clone(), + key, + )?; + context.send_by_key(packet.buffer(), route_key)?; + self.rsa_cipher.lock().replace(rsa_cipher); + } + return Ok(()); + } + + let handshake_info = HandshakeInfo::new_no_secret(response.version); + if self.callback.handshake(handshake_info) { + //没有加密,则发送注册请求 + self.register(current_device, context)?; + } + + return Ok(()); + } + //服务端数据解密 + self.server_cipher.decrypt_ipv4(&mut net_packet)?; + match net_packet.protocol() { + Protocol::Service => { + self.service(context, current_device, net_packet, route_key)?; + } + Protocol::Error => { + self.error(context, current_device, net_packet, route_key)?; + } + Protocol::Control => { + self.control(context, current_device, net_packet, route_key)?; + } + Protocol::IpTurn => { + match ip_turn_packet::Protocol::from(net_packet.transport_protocol()) { + ip_turn_packet::Protocol::Ipv4 => { + let ipv4 = IpV4Packet::new(net_packet.payload())?; + match ipv4.protocol() { + ipv4::protocol::Protocol::Icmp => { + if ipv4.destination_ip() == current_device.virtual_ip { + let icmp_packet = icmp::IcmpPacket::new(ipv4.payload())?; + if icmp_packet.kind() == Kind::EchoReply { + //网关ip ping的回应 + self.device.write(net_packet.payload())?; + return Ok(()); + } + } + } + _ => {} + } + } + ip_turn_packet::Protocol::Ipv4Broadcast => {} + ip_turn_packet::Protocol::Unknown(_) => {} + } + } + Protocol::OtherTurn => {} + Protocol::Unknown(_) => {} + } + Ok(()) + } +} + +impl ServerPacketHandler { + fn service( + &self, + context: &Context, + current_device: &CurrentDeviceInfo, + net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + ) -> io::Result<()> { + match service_packet::Protocol::from(net_packet.transport_protocol()) { + service_packet::Protocol::RegistrationResponse => { + let response = RegistrationResponse::parse_from_bytes(net_packet.payload()) + .map_err(|e| { + io::Error::new( + io::ErrorKind::Other, + format!("RegistrationResponse {:?}", e), + ) + })?; + let virtual_ip = Ipv4Addr::from(response.virtual_ip); + let virtual_netmask = Ipv4Addr::from(response.virtual_netmask); + let virtual_gateway = Ipv4Addr::from(response.virtual_gateway); + let virtual_network = + Ipv4Addr::from(response.virtual_ip & response.virtual_netmask); + let register_info = RegisterInfo::new(virtual_ip, virtual_netmask, virtual_gateway); + if self.callback.register(register_info) { + let route = Route::from(route_key, 1, 199); + context + .route_table + .add_route_if_absent(virtual_gateway, route); + let old = current_device; + let mut cur = *current_device; + loop { + let mut new_current_device = cur; + new_current_device.update(virtual_ip, virtual_netmask, virtual_gateway); + new_current_device.virtual_ip = virtual_ip; + new_current_device.virtual_netmask = virtual_netmask; + new_current_device.virtual_gateway = virtual_gateway; + if let Err(c) = self + .current_device + .compare_exchange(cur, new_current_device) + { + cur = c; + } else { + break; + } + } + let _ = context.change_status(&self.current_device); + + let public_ip = response.public_ip.into(); + let public_port = response.public_port as u16; + self.nat_test + .update_addr(route_key.index(), public_ip, public_port); + if old.virtual_ip != virtual_ip + || old.virtual_gateway != virtual_gateway + || old.virtual_netmask != virtual_netmask + { + if old.virtual_ip != Ipv4Addr::UNSPECIFIED { + log::info!("ip发生变化,old:{:?},response={:?}", old, response); + } + self.device.set_ip(virtual_ip, virtual_netmask)?; + let mut guard = self.route_record.lock(); + for (dest, mask) in guard.drain(..) { + if let Err(e) = self.device.delete_route(dest, mask) { + log::warn!("删除路由失败 ={:?}", e); + } + } + if let Err(e) = self.device.add_route(virtual_network, virtual_netmask, 1) { + log::warn!("添加默认路由失败 ={:?}", e); + } else { + guard.push((virtual_network, virtual_netmask)); + } + for (dest, mask) in self.external_route.to_route() { + if let Err(e) = self.device.add_route(dest, mask, 1) { + log::warn!("添加路由失败 ={:?}", e); + } else { + guard.push((dest, mask)); + } + } + } + self.set_device_info_list(response.device_info_list, response.epoch as _); + } + } + service_packet::Protocol::RegistrationRequest => { + //不处理注册包 + } + service_packet::Protocol::PollDeviceList => {} + service_packet::Protocol::PushDeviceList => { + let response = DeviceList::parse_from_bytes(net_packet.payload()).map_err(|e| { + io::Error::new(io::ErrorKind::Other, format!("PushDeviceList {:?}", e)) + })?; + self.set_device_info_list(response.device_info_list, response.epoch as _); + } + service_packet::Protocol::HandshakeRequest => {} + service_packet::Protocol::HandshakeResponse => {} + service_packet::Protocol::SecretHandshakeRequest => {} + service_packet::Protocol::SecretHandshakeResponse => { + //加密握手结束,发送注册数据 + self.register(current_device, context)?; + } + service_packet::Protocol::Unknown(e) => { + log::warn!("service_packet::Protocol::Unknown = {}", e); + } + } + Ok(()) + } + fn set_device_info_list(&self, device_info_list: Vec, epoch: u16) { + let ip_list: Vec = device_info_list + .into_iter() + .map(|info| { + PeerDeviceInfo::new( + Ipv4Addr::from(info.virtual_ip), + info.name, + info.device_status as u8, + info.client_secret, + ) + }) + .collect(); + let mut dev = self.device_list.lock(); + //这里可能会收到旧的消息,但是随着时间推移总会收到新的 + dev.0 = epoch; + dev.1 = ip_list; + } + fn register(&self, current_device: &CurrentDeviceInfo, context: &Context) -> io::Result<()> { + if current_device.status == ConnectStatus::Connected { + //已连接的不需要注册 + return Ok(()); + } + let token = self.config_info.token.clone(); + let device_id = self.config_info.device_id.clone(); + let name = self.config_info.name.clone(); + let client_secret = self.config_info.client_secret; + let ip = self.config_info.ip; + let response = registrar::registration_request_packet( + &self.server_cipher, + token, + device_id, + name, + ip, + false, + false, + client_secret, + )?; + //注册请求只发送到默认通道 + context.send_default(response.buffer(), current_device.connect_server) + } + fn error( + &self, + _context: &Context, + _current_device: &CurrentDeviceInfo, + net_packet: NetPacket<&mut [u8]>, + _route_key: RouteKey, + ) -> io::Result<()> { + match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload())? { + InErrorPacket::TokenError => { + // token错误,可能是服务端设置了白名单 + let err = ErrorInfo::new(ErrorType::TokenError); + self.callback.error(err); + } + InErrorPacket::Disconnect => { + let err = ErrorInfo::new(ErrorType::Disconnect); + self.callback.error(err); + //掉线epoch要归零 + { + let mut dev = self.device_list.lock(); + dev.0 = 0; + drop(dev); + } + // self.register(current_device, context, route_key)?; + } + InErrorPacket::AddressExhausted => { + // 地址用尽 + let err = ErrorInfo::new(ErrorType::AddressExhausted); + self.callback.error(err); + } + InErrorPacket::OtherError(e) => { + let err = ErrorInfo::new_msg(ErrorType::Unknown, e.message()?); + self.callback.error(err); + } + InErrorPacket::IpAlreadyExists => { + log::error!("IpAlreadyExists"); + let err = ErrorInfo::new(ErrorType::IpAlreadyExists); + self.callback.error(err); + } + InErrorPacket::InvalidIp => { + log::error!("InvalidIp"); + let err = ErrorInfo::new(ErrorType::InvalidIp); + self.callback.error(err); + } + InErrorPacket::NoKey => { + //这个类型最开头已经处理过,这里忽略 + } + } + Ok(()) + } + fn control( + &self, + context: &Context, + current_device: &CurrentDeviceInfo, + net_packet: NetPacket<&mut [u8]>, + route_key: RouteKey, + ) -> io::Result<()> { + match ControlPacket::new(net_packet.transport_protocol(), net_packet.payload())? { + ControlPacket::PongPacket(pong_packet) => { + let current_time = crate::handle::now_time() as u16; + if current_time < pong_packet.time() { + return Ok(()); + } + let metric = net_packet.source_ttl() - net_packet.ttl() + 1; + let rt = (current_time - pong_packet.time()) as i64; + let route = Route::from(route_key, metric, rt); + context.route_table.add_route(net_packet.source(), route); + let epoch = self.device_list.lock().0; + if pong_packet.epoch() != epoch { + //纪元不一致,可能有新客户端连接,向服务端拉取客户端列表 + let mut poll_device = NetPacket::new_encrypt([0; 12 + ENCRYPTION_RESERVED])?; + poll_device.set_source(current_device.virtual_ip); + poll_device.set_destination(GATEWAY_IP); + poll_device.set_version(Version::V1); + poll_device.set_gateway_flag(true); + poll_device.first_set_ttl(MAX_TTL); + poll_device.set_protocol(Protocol::Service); + poll_device + .set_transport_protocol(service_packet::Protocol::PollDeviceList.into()); + self.server_cipher.encrypt_ipv4(&mut poll_device)?; + //发送到默认服务端即可 + context.send_default(poll_device.buffer(), current_device.connect_server)?; + } + } + ControlPacket::AddrResponse(addr_packet) => { + //更新本地公网ipv4 + self.nat_test.update_addr( + route_key.index(), + addr_packet.ipv4(), + addr_packet.port(), + ); + } + _ => {} + } + Ok(()) + } +} diff --git a/vnt/src/handle/recv_data/turn.rs b/vnt/src/handle/recv_data/turn.rs new file mode 100644 index 0000000..508676d --- /dev/null +++ b/vnt/src/handle/recv_data/turn.rs @@ -0,0 +1,39 @@ +use crate::channel::context::Context; +use crate::channel::RouteKey; +use crate::handle::recv_data::PacketHandler; +use crate::handle::CurrentDeviceInfo; +use crate::protocol::NetPacket; + +/// 处理客户端中转包 +#[derive(Clone)] +pub struct TurnPacketHandler {} + +impl TurnPacketHandler { + pub fn new() -> Self { + Self {} + } +} + +impl PacketHandler for TurnPacketHandler { + fn handle( + &self, + mut net_packet: NetPacket<&mut [u8]>, + _route_key: RouteKey, + context: &Context, + _current_device: &CurrentDeviceInfo, + ) -> std::io::Result<()> { + // ttl减一 + let ttl = net_packet.incr_ttl(); + if ttl > 0 { + let destination = net_packet.destination(); + if let Some(route) = context.route_table.route_one(&destination) { + if route.metric <= ttl { + context.send_by_key(net_packet.buffer(), route.route_key())?; + } + } + //其他没有路由的不转发 + } + + Ok(()) + } +} diff --git a/vnt/src/handle/registrar.rs b/vnt/src/handle/registrar.rs new file mode 100644 index 0000000..b230923 --- /dev/null +++ b/vnt/src/handle/registrar.rs @@ -0,0 +1,49 @@ +use std::io; +use std::net::Ipv4Addr; + +use protobuf::Message; + +use crate::cipher::Cipher; +use crate::handle::{GATEWAY_IP, SELF_IP}; +use crate::proto::message::RegistrationRequest; +use crate::protocol::body::ENCRYPTION_RESERVED; +use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL}; + +/// 注册数据 +pub fn registration_request_packet( + server_cipher: &Cipher, + token: String, + device_id: String, + name: String, + ip: Option, + is_fast: bool, + allow_ip_change: bool, + client_secret: bool, +) -> io::Result>> { + let mut request = RegistrationRequest::new(); + request.token = token; + request.device_id = device_id; + request.name = name; + if let Some(ip) = ip { + request.virtual_ip = ip.into(); + } + request.allow_ip_change = allow_ip_change; + request.is_fast = is_fast; + request.version = crate::VNT_VERSION.to_string(); + request.client_secret = client_secret; + let bytes = request.write_to_bytes().map_err(|e| { + io::Error::new(io::ErrorKind::Other, format!("RegistrationRequest {:?}", e)) + })?; + let buf = vec![0u8; 12 + bytes.len() + ENCRYPTION_RESERVED]; + let mut net_packet = NetPacket::new_encrypt(buf)?; + net_packet.set_destination(GATEWAY_IP); + net_packet.set_source(SELF_IP); + net_packet.set_version(Version::V1); + net_packet.set_gateway_flag(true); + net_packet.set_protocol(Protocol::Service); + net_packet.set_transport_protocol(service_packet::Protocol::RegistrationRequest.into()); + net_packet.first_set_ttl(MAX_TTL); + net_packet.set_payload(&bytes)?; + server_cipher.encrypt_ipv4(&mut net_packet)?; + Ok(net_packet) +} diff --git a/vnt/src/handle/tun_tap/channel_group.rs b/vnt/src/handle/tun_tap/channel_group.rs index 10160e0..531a413 100644 --- a/vnt/src/handle/tun_tap/channel_group.rs +++ b/vnt/src/handle/tun_tap/channel_group.rs @@ -1,30 +1,30 @@ -#[derive(Clone)] -pub struct BufSenderGroup( - usize, - Vec, usize, usize)>>, -); +use std::sync::mpsc::{sync_channel, Receiver, SendError, SyncSender}; -pub struct BufReceiverGroup(pub Vec, usize, usize)>>); - -impl BufSenderGroup { - pub fn send(&mut self, val: (Vec, usize, usize)) -> bool { - let index = self.0 % self.1.len(); - self.0 = self.0.wrapping_add(1); - self.1[index].send(val).is_ok() - } -} - -pub fn buf_channel_group(size: usize) -> (BufSenderGroup, BufReceiverGroup) { - let mut buf_sender_group = Vec::with_capacity(size); - let mut buf_receiver_group = Vec::with_capacity(size); +pub fn channel_group(size: usize, bound: usize) -> (GroupSyncSender, Vec>) { + let mut senders = Vec::with_capacity(size); + let mut receivers = Vec::with_capacity(size); for _ in 0..size { - let (buf_sender, buf_receiver) = - std::sync::mpsc::sync_channel::<(Vec, usize, usize)>(1); - buf_sender_group.push(buf_sender); - buf_receiver_group.push(buf_receiver); + let (s, r) = sync_channel(bound); + senders.push(s); + receivers.push(r); } ( - BufSenderGroup(0, buf_sender_group), - BufReceiverGroup(buf_receiver_group), + GroupSyncSender { + count: 0, + base: senders, + }, + receivers, ) } + +pub struct GroupSyncSender { + count: usize, + base: Vec>, +} + +impl GroupSyncSender { + pub fn send(&mut self, t: T) -> Result<(), SendError> { + self.count += 1; + self.base[self.count % self.base.len()].send(t) + } +} diff --git a/vnt/src/handle/tun_tap/mod.rs b/vnt/src/handle/tun_tap/mod.rs index 8ed1fc6..cfbbb11 100644 --- a/vnt/src/handle/tun_tap/mod.rs +++ b/vnt/src/handle/tun_tap/mod.rs @@ -1,17 +1,13 @@ +use std::io; use std::net::Ipv4Addr; -use std::sync::Arc; - -use parking_lot::RwLock; +use crate::channel::context::Context; use packet::ip::ipv4::packet::IpV4Packet; use packet::ip::ipv4::protocol::Protocol; -use crate::channel::sender::ChannelSender; use crate::cipher::Cipher; -use crate::error::*; use crate::external_route::ExternalRoute; use crate::handle::{check_dest, CurrentDeviceInfo}; -use crate::igmp_server::{IgmpServer, Multicast}; #[cfg(feature = "ip_proxy")] use crate::ip_proxy::{IpProxyMap, ProxyHandler}; use crate::protocol; @@ -19,20 +15,17 @@ use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::ip_turn_packet::BroadcastPacket; use crate::protocol::{ip_turn_packet, NetPacket, Version, MAX_TTL}; -pub mod channel_group; -#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] -pub mod tap_handler; +mod channel_group; pub mod tun_handler; fn broadcast( server_cipher: &Cipher, - multicast_members: Option>>, - sender: &ChannelSender, + sender: &Context, net_packet: &mut NetPacket<&mut [u8]>, current_device: &CurrentDeviceInfo, -) -> Result<()> { +) -> io::Result<()> { let mut peer_ips = Vec::with_capacity(8); - let vec = sender.route_table_one(); + let vec = sender.route_table.route_table_one(); let mut relay_count = 0; const MAX_COUNT: usize = 8; for (peer_ip, route) in vec { @@ -42,14 +35,9 @@ fn broadcast( if peer_ips.len() == MAX_COUNT { break; } - if let Some(members) = &multicast_members { - if !members.read().is_send(&peer_ip) { - continue; - } - } if route.is_p2p() && sender - .try_send_by_key(net_packet.buffer(), &route.route_key()) + .send_by_key(net_packet.buffer(), route.route_key()) .is_ok() { peer_ips.push(peer_ip); @@ -63,7 +51,7 @@ fn broadcast( } //转发到服务端的可选择广播,还要进行服务端加密 if peer_ips.is_empty() { - sender.send_main(net_packet.buffer(), current_device.connect_server)?; + sender.send_default(net_packet.buffer(), current_device.connect_server)?; } else { let buf = vec![0u8; 12 + 1 + peer_ips.len() * 4 + net_packet.data_len() + ENCRYPTION_RESERVED]; @@ -82,7 +70,7 @@ fn broadcast( broadcast.set_address(&peer_ips)?; broadcast.set_data(net_packet.buffer())?; server_cipher.encrypt_ipv4(&mut server_packet)?; - sender.send_main(server_packet.buffer(), current_device.connect_server)?; + sender.send_default(server_packet.buffer(), current_device.connect_server)?; } Ok(()) } @@ -92,16 +80,15 @@ fn broadcast( /// #[inline] pub fn base_handle( - sender: &ChannelSender, + context: &Context, buf: &mut [u8], data_len: usize, //数据总长度=12+ip包长度 - igmp_server: &Option, current_device: CurrentDeviceInfo, - ip_route: &Option, + ip_route: &ExternalRoute, #[cfg(feature = "ip_proxy")] proxy_map: &Option, client_cipher: &Cipher, server_cipher: &Cipher, -) -> Result<()> { +) -> io::Result<()> { let ipv4_packet = IpV4Packet::new(&buf[12..data_len])?; let protocol = ipv4_packet.protocol(); let src_ip = ipv4_packet.source_ip(); @@ -117,52 +104,26 @@ pub fn base_handle( if protocol == Protocol::Icmp { net_packet.set_gateway_flag(true); server_cipher.encrypt_ipv4(&mut net_packet)?; - sender.send_main(net_packet.buffer(), current_device.connect_server)?; + context.send_default(net_packet.buffer(), current_device.connect_server)?; } return Ok(()); } if dest_ip.is_multicast() { match protocol { - Protocol::Igmp => { - if igmp_server.is_some() { - //发送到服务端 - net_packet.set_destination(current_device.virtual_gateway); - net_packet.set_gateway_flag(true); - server_cipher.encrypt_ipv4(&mut net_packet)?; - sender.send_main(net_packet.buffer(), current_device.connect_server)?; - } - } Protocol::Udp => { - let multicast_members = if let Some(igmp_server) = igmp_server { - igmp_server.load(&dest_ip) - } else { - //当作广播处理 - net_packet.set_destination(Ipv4Addr::BROADCAST); - None - }; + //当作广播处理 + net_packet.set_destination(Ipv4Addr::BROADCAST); client_cipher.encrypt_ipv4(&mut net_packet)?; - broadcast( - server_cipher, - multicast_members, - sender, - &mut net_packet, - ¤t_device, - )?; + broadcast(server_cipher, context, &mut net_packet, ¤t_device)?; } _ => {} } return Ok(()); } - if dest_ip.is_broadcast() || current_device.broadcast_address == dest_ip { + if dest_ip.is_broadcast() || current_device.broadcast_ip == dest_ip { // 广播 发送到直连目标 client_cipher.encrypt_ipv4(&mut net_packet)?; - broadcast( - server_cipher, - None, - sender, - &mut net_packet, - ¤t_device, - )?; + broadcast(server_cipher, context, &mut net_packet, ¤t_device)?; return Ok(()); } if !check_dest( @@ -170,18 +131,14 @@ pub fn base_handle( current_device.virtual_netmask, current_device.virtual_network, ) { - if let Some(ip_route) = ip_route { - if let Some(r_dest_ip) = ip_route.route(&dest_ip) { - //路由的目标不能是自己 - if r_dest_ip == src_ip { - return Ok(()); - } - //需要修改目的地址 - dest_ip = r_dest_ip; - net_packet.set_destination(r_dest_ip); - } else { + if let Some(r_dest_ip) = ip_route.route(&dest_ip) { + //路由的目标不能是自己 + if r_dest_ip == src_ip { return Ok(()); } + //需要修改目的地址 + dest_ip = r_dest_ip; + net_packet.set_destination(r_dest_ip); } else { return Ok(()); } @@ -193,11 +150,8 @@ pub fn base_handle( } client_cipher.encrypt_ipv4(&mut net_packet)?; //优先发到直连到地址 - if sender - .try_send_by_id(net_packet.buffer(), &dest_ip) - .is_err() - { - sender.send_main(net_packet.buffer(), current_device.connect_server)?; + if context.send_by_id(net_packet.buffer(), &dest_ip).is_err() { + context.send_default(net_packet.buffer(), current_device.connect_server)?; } return Ok(()); } diff --git a/vnt/src/handle/tun_tap/tun_handler.rs b/vnt/src/handle/tun_tap/tun_handler.rs index 257d07f..fe03108 100644 --- a/vnt/src/handle/tun_tap/tun_handler.rs +++ b/vnt/src/handle/tun_tap/tun_handler.rs @@ -1,25 +1,26 @@ +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; -use crate::channel::sender::ChannelSender; -use crate::cipher::Cipher; -use crate::core::status::VntWorker; use packet::icmp::icmp::IcmpPacket; use packet::icmp::Kind; use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; +use tun::device::IFace; +use tun::Device; -use crate::error::*; +use crate::channel::context::Context; +use crate::cipher::Cipher; use crate::external_route::ExternalRoute; -use crate::handle::tun_tap::channel_group::{buf_channel_group, BufSenderGroup}; +use crate::handle::tun_tap::channel_group::{channel_group, GroupSyncSender}; use crate::handle::CurrentDeviceInfo; -use crate::igmp_server::IgmpServer; #[cfg(feature = "ip_proxy")] use crate::ip_proxy::IpProxyMap; -use crate::tun_tap_device::{DeviceReader, DeviceWriter}; -fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> Result<()> { +use crate::util::StopManager; + +fn icmp(device_writer: &Device, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> io::Result<()> { if ipv4_packet.protocol() == ipv4::protocol::Protocol::Icmp { let mut icmp = IcmpPacket::new(ipv4_packet.payload_mut())?; if icmp.kind() == Kind::EchoRequest { @@ -29,26 +30,28 @@ fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> ipv4_packet.set_source_ip(ipv4_packet.destination_ip()); ipv4_packet.set_destination_ip(src); ipv4_packet.update_checksum(); - device_writer.write_ipv4_tun(ipv4_packet.buffer)?; + device_writer.write(ipv4_packet.buffer)?; } } Ok(()) } /// 接收tun数据,并且转发到udp上 -#[inline] fn handle( - sender: &ChannelSender, + context: &Context, data: &mut [u8], len: usize, - device_writer: &DeviceWriter, - igmp_server: &Option, + device_writer: &Device, current_device: CurrentDeviceInfo, - ip_route: &Option, + ip_route: &ExternalRoute, #[cfg(feature = "ip_proxy")] proxy_map: &Option, client_cipher: &Cipher, server_cipher: &Cipher, -) -> Result<()> { +) -> io::Result<()> { + if len > 12 && data[12] >> 4 != 4 { + //忽略非ipv4包 + return Ok(()); + } let ipv4_packet = IpV4Packet::new(&mut data[12..len])?; let src_ip = ipv4_packet.source_ip(); let dest_ip = ipv4_packet.destination_ip(); @@ -56,10 +59,9 @@ fn handle( return icmp(&device_writer, ipv4_packet); } return crate::handle::tun_tap::base_handle( - sender, + context, data, len, - igmp_server, current_device, ip_route, #[cfg(feature = "ip_proxy")] @@ -70,143 +72,127 @@ fn handle( } pub fn start( - worker: VntWorker, - sender: ChannelSender, - device_reader: DeviceReader, - device_writer: DeviceWriter, - igmp_server: Option, + stop_manager: StopManager, + context: Context, + device: Arc, current_device: Arc>, - ip_route: Option, + ip_route: ExternalRoute, #[cfg(feature = "ip_proxy")] ip_proxy_map: Option, client_cipher: Cipher, server_cipher: Cipher, parallel: usize, -) { - if parallel == 1 { - thread::Builder::new() - .name("tun_handler".into()) - .spawn(move || { - if let Err(e) = start_simple( - &sender, - device_reader, - &device_writer, - igmp_server, - current_device, - ip_route, - #[cfg(feature = "ip_proxy")] - ip_proxy_map, - client_cipher, - server_cipher, - ) { - log::warn!("stop:{}", e); - } - let _ = sender.close(); - let _ = device_writer.close(); - worker.stop_all(); - }) - .unwrap(); - } else { - let (buf_sender, buf_receiver) = buf_channel_group(parallel); - for buf_receiver in buf_receiver.0 { - let sender = sender.clone(); - let device_writer = device_writer.clone(); - let igmp_server = igmp_server.clone(); + up_counter: Arc, +) -> io::Result<()> { + let worker = { + let device = device.clone(); + stop_manager.add_listener("tun_device".into(), move || { + if let Err(e) = device.shutdown() { + log::warn!("{:?}", e); + } + })? + }; + if parallel > 1 { + let (sender, receivers) = channel_group::<(Vec, usize)>(parallel, 16); + for (index, receiver) in receivers.into_iter().enumerate() { + let context = context.clone(); + let device = device.clone(); let current_device = current_device.clone(); let ip_route = ip_route.clone(); #[cfg(feature = "ip_proxy")] let ip_proxy_map = ip_proxy_map.clone(); let client_cipher = client_cipher.clone(); let server_cipher = server_cipher.clone(); - thread::spawn(move || { - while let Ok((mut buf, start, len)) = buf_receiver.recv() { - match handle( - &sender, - &mut buf[start..], - len, - &device_writer, - &igmp_server, - current_device.load(), - &ip_route, - #[cfg(feature = "ip_proxy")] - &ip_proxy_map, - &client_cipher, - &server_cipher, - ) { - Ok(_) => {} - Err(e) => { - log::warn!("{:?}", e) + thread::Builder::new() + .name(format!("tun_handler_{}", index)) + .spawn(move || { + while let Ok((mut buf, len)) = receiver.recv() { + #[cfg(not(target_os = "macos"))] + let start = 0; + #[cfg(target_os = "macos")] + let start = 4; + match handle( + &context, + &mut buf[start..], + len, + &device, + current_device.load(), + &ip_route, + #[cfg(feature = "ip_proxy")] + &ip_proxy_map, + &client_cipher, + &server_cipher, + ) { + Ok(_) => {} + Err(e) => { + log::warn!("{:?}", e) + } } } - } - let _ = sender.close(); - let _ = device_writer.close(); - }); + })?; } - thread::Builder::new() .name("tun_handler".into()) .spawn(move || { - if let Err(e) = start_(&sender, device_reader, buf_sender) { + if let Err(e) = start_multi(stop_manager, device, sender, &up_counter) { log::warn!("stop:{}", e); } - let _ = sender.close(); - let _ = device_writer.close(); worker.stop_all(); - }) - .unwrap(); - } -} - -fn start_( - sender: &ChannelSender, - device_reader: DeviceReader, - mut buf_sender: BufSenderGroup, -) -> io::Result<()> { - loop { - let mut buf = vec![0; 4096]; - buf[..12].fill(0); - if sender.is_close() { - return Ok(()); - } - let start = 0; - let len = device_reader.read(&mut buf[12..])? + 12; - #[cfg(any(target_os = "macos"))] - let start = 4; - if !buf_sender.send((buf, start, len)) { - return Err(io::Error::new( - io::ErrorKind::Other, - "tun buf_sender发送失败", - )); - } + })?; + } else { + thread::Builder::new() + .name("tun_handler".into()) + .spawn(move || { + if let Err(e) = start_simple( + stop_manager, + &context, + device, + current_device, + ip_route, + #[cfg(feature = "ip_proxy")] + ip_proxy_map, + client_cipher, + server_cipher, + &up_counter, + ) { + log::warn!("stop:{}", e); + } + worker.stop_all(); + })?; } + Ok(()) } fn start_simple( - sender: &ChannelSender, - device_reader: DeviceReader, - device_writer: &DeviceWriter, - igmp_server: Option, + stop_manager: StopManager, + context: &Context, + device: Arc, current_device: Arc>, - ip_route: Option, + ip_route: ExternalRoute, #[cfg(feature = "ip_proxy")] ip_proxy_map: Option, client_cipher: Cipher, server_cipher: Cipher, + up_counter: &AtomicU64, ) -> io::Result<()> { - let mut buf = [0; 4096]; + let mut buf = [0; 1024 * 16]; loop { - if sender.is_close() { + if stop_manager.is_stop() { return Ok(()); } - buf[..12].fill(0); - let len = device_reader.read(&mut buf[12..])? + 12; + let len = device.read(&mut buf[12..])? + 12; + //单线程的 + up_counter.store( + up_counter.load(Ordering::Relaxed) + len as u64, + Ordering::Relaxed, + ); #[cfg(any(target_os = "macos"))] let mut buf = &mut buf[4..]; + // buf是重复利用的,需要重置头部 + buf[..12].fill(0); match handle( - sender, + context, &mut buf, len, - device_writer, - &igmp_server, + &device, current_device.load(), &ip_route, #[cfg(feature = "ip_proxy")] @@ -221,3 +207,26 @@ fn start_simple( } } } + +fn start_multi( + stop_manager: StopManager, + device: Arc, + mut group_sync_sender: GroupSyncSender<(Vec, usize)>, + up_counter: &AtomicU64, +) -> io::Result<()> { + loop { + if stop_manager.is_stop() { + return Ok(()); + } + let mut buf = vec![0; 1024 * 16]; + let len = device.read(&mut buf[12..])? + 12; + //单线程的 + up_counter.store( + up_counter.load(Ordering::Relaxed) + len as u64, + Ordering::Relaxed, + ); + if group_sync_sender.send((buf, len)).is_err() { + return Ok(()); + } + } +}