diff --git a/switch/Cargo.toml b/switch/Cargo.toml index 6ee23b9..fb2b1cb 100644 --- a/switch/Cargo.toml +++ b/switch/Cargo.toml @@ -13,21 +13,15 @@ libc = "0.2.137" crossbeam-utils = "0.8" crossbeam-skiplist = "0.1" parking_lot = "0.12.1" - -#rsa = "0.7.2" rand = "0.8.5" sha2 = { version = "0.10.6", features = ["oid"] } aes-gcm = "0.10.2" thiserror = "1.0.37" -#chrono = "0.4.23" -#lazy_static = "1.4.0" -#moka = "0.9.6" protobuf = "3.2.0" -#local-ip-address = "0.4.9" socket2 ={ version = "0.5.2", features = ["all"] } tokio = { version = "1.28.1", features = ["full"] } -[target.'cfg(any(unix))'.dependencies] +[target.'cfg(any(target_os = "linux",target_os = "macos"))'.dependencies] tun = { path = "./rust-tun" } [target.'cfg(target_os = "windows")'.dependencies] diff --git a/switch/rust-tun/src/platform/mod.rs b/switch/rust-tun/src/platform/mod.rs index 2ecf9e1..cc5b674 100644 --- a/switch/rust-tun/src/platform/mod.rs +++ b/switch/rust-tun/src/platform/mod.rs @@ -27,15 +27,6 @@ pub mod macos; #[cfg(target_os = "macos")] pub use self::macos::{create, Configuration, Device, Queue}; -#[cfg(target_os = "ios")] -pub mod ios; -#[cfg(target_os = "ios")] -pub use self::ios::{create, Configuration, Device, Queue}; - -#[cfg(target_os = "android")] -pub mod android; -#[cfg(target_os = "android")] -pub use self::android::{create, Configuration, Device, Queue}; #[cfg(test)] mod test { diff --git a/switch/src/channel/channel.rs b/switch/src/channel/channel.rs index 0075edf..73ebf23 100644 --- a/switch/src/channel/channel.rs +++ b/switch/src/channel/channel.rs @@ -10,10 +10,11 @@ use tokio::net::UdpSocket; use tokio::sync::watch::{channel, Receiver, Sender}; use crate::channel::{Route, RouteKey, Status}; use crate::channel::punch::NatType; +use crate::core::status::SwitchWorker; use crate::handle::recv_handler::ChannelDataHandler; pub struct ContextInner { - pub(crate) lock:Mutex<()>, + pub(crate) lock: Mutex<()>, pub(crate) count: AtomicUsize, pub(crate) main_channel: Arc, pub(crate) route_table: SkipMap>, @@ -35,7 +36,7 @@ impl Context { let channel_num = 1; let (status_sender, status_receiver) = channel(Status::Cone); let inner = Arc::new(ContextInner { - lock:Mutex::new(()), + lock: Mutex::new(()), count: AtomicUsize::new(0), main_channel, route_table: SkipMap::new(), @@ -143,9 +144,9 @@ impl Context { fn add_route_(&self, id: Ipv4Addr, route: Route, only_if_absent: bool) { let key = route.route_key(); let guard = self.inner.lock.lock(); - let mut list = if let Some(entry) = self.inner.route_table.get(&id){ + let mut list = if let Some(entry) = self.inner.route_table.get(&id) { entry.value().clone() - }else{ + } else { Vec::with_capacity(4) }; let mut exist = false; @@ -178,7 +179,7 @@ impl Context { list.truncate(max_len); } } - self.inner.route_table.insert(id,list); + self.inner.route_table.insert(id, list); self.inner.route_table_time.insert((key, id), AtomicCell::new(Instant::now())); drop(guard); } @@ -250,7 +251,7 @@ impl Context { let mut routes = v.value().clone(); drop(v); routes.retain(|x| x.route_key() != route_key); - self.inner.route_table.insert(*id,routes); + self.inner.route_table.insert(*id, routes); self.inner.route_table_time.remove(&(route_key, *id)); } drop(guard); @@ -294,54 +295,63 @@ impl Channel { } } pub async fn start(self, + mut worker: SwitchWorker, head_reserve: usize,//头部预留字节 symmetric_channel_num: usize,//对称网络,则再加一组监听,提升打洞成功率 ) { let context = self.context; let main_channel = context.inner.main_channel.clone(); let handler = self.handler.clone(); - tokio::spawn(Self::start_(context.clone(), handler.clone(), main_channel.clone(), head_reserve, true)); - tokio::spawn(Self::start_(context.clone(), handler, main_channel, head_reserve, true)); + tokio::spawn(Self::start_(worker.clone(), context.clone(), handler.clone(), main_channel.clone(), head_reserve, true)); + tokio::spawn(Self::start_(worker.clone(), context.clone(), handler, main_channel, head_reserve, true)); let mut cur_status = Status::Cone; let mut status_receiver = context.inner.status_receiver.clone(); loop { - match status_receiver.changed().await { - Ok(_) => { - match *status_receiver.borrow() { - Status::Cone => { - cur_status = Status::Cone; - } - Status::Symmetric => { - if cur_status == Status::Symmetric { - continue; - } - cur_status = Status::Symmetric; - for _ in 0..symmetric_channel_num { - match UdpSocket::bind("0.0.0.0:0").await { - Ok(udp) => { - let udp = Arc::new(udp); - let context = context.clone(); - let handler = self.handler.clone(); - tokio::spawn(Self::start_(context, handler, udp, head_reserve, false)); + tokio::select! { + _=worker.stop_wait()=>{ + break; + } + rs=status_receiver.changed()=>{ + match rs { + Ok(_) => { + match *status_receiver.borrow() { + Status::Cone => { + cur_status = Status::Cone; + } + Status::Symmetric => { + if cur_status == Status::Symmetric { + continue; } - Err(e) => { - log::error!("{}",e); + cur_status = Status::Symmetric; + for _ in 0..symmetric_channel_num { + match UdpSocket::bind("0.0.0.0:0").await { + Ok(udp) => { + let udp = Arc::new(udp); + let context = context.clone(); + let handler = self.handler.clone(); + tokio::spawn(Self::start_(worker.clone(),context, handler, udp, head_reserve, false)); + } + Err(e) => { + log::error!("{}",e); + } + } } } + Status::Close => { + break; + } } } - Status::Close => { + Err(_) => { break; } } } - Err(_) => { - break; - } } } + worker.stop_all(); } - async fn start_(context: Context, + async fn start_(mut worker: SwitchWorker, context: Context, mut handler: ChannelDataHandler, udp: Arc, head_reserve: usize, @@ -382,8 +392,14 @@ impl Channel { } } } + _=worker.stop_wait()=>{ + break; + } } } context.inner.udp_map.remove(&id); + if is_core { + worker.stop_all(); + } } } diff --git a/switch/src/core/mod.rs b/switch/src/core/mod.rs index cf4c831..2a5b36d 100644 --- a/switch/src/core/mod.rs +++ b/switch/src/core/mod.rs @@ -2,34 +2,44 @@ use std::{io, thread}; use std::net::{Ipv4Addr, SocketAddr}; use std::sync::Arc; use std::time::Duration; -use aes_gcm::{Aes256Gcm, Key, KeyInit}; -use crossbeam_utils::atomic::AtomicCell; +use aes_gcm::{Aes256Gcm, Key, KeyInit}; use crossbeam_skiplist::SkipMap; +use crossbeam_utils::atomic::AtomicCell; use parking_lot::Mutex; +use sha2::Digest; use tokio::net::UdpSocket; use tokio::sync::mpsc::channel; - +use crate::channel::{Route, RouteKey}; use crate::channel::channel::{Channel, Context}; use crate::channel::idle::Idle; use crate::channel::punch::{NatInfo, Punch}; -use crate::channel::{Route, RouteKey}; use crate::channel::sender::ChannelSender; - +use crate::core::status::SwitchStatusManger; +use crate::error::Error; use crate::external_route::ExternalRoute; use crate::handle::{ConnectStatus, CurrentDeviceInfo, heartbeat_handler, PeerDeviceInfo, punch_handler, registration_handler}; use crate::handle::recv_handler::ChannelDataHandler; -use crate::handle::tun_tap::{tap_handler, tun_handler}; +use crate::handle::registration_handler::{RegResponse, ReqEnum}; +#[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] +use crate::handle::tun_tap::tap_handler; +use crate::handle::tun_tap::tun_handler; use crate::igmp_server::IgmpServer; use crate::nat::NatTest; use crate::tun_tap_device; -use crate::tun_tap_device::DeviceWriter; +use crate::tun_tap_device::{DeviceReader, DeviceWriter}; + +pub mod status; +pub mod sync; + pub struct Switch { name: String, current_device: Arc>, context: Context, + switch_status_manager: SwitchStatusManger, + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] device_writer: DeviceWriter, /// 0. 机器纪元,每一次上线或者下线都会增1,用于感知网络中机器变化 /// 服务端和客户端的不一致,则服务端会推送新的设备列表 @@ -40,37 +50,121 @@ pub struct Switch { peer_nat_info_map: Arc>, } -impl Switch { - pub async fn start(config: Config) -> crate::Result { - log::info!("config:{:?}",config); +pub struct SwitchUtil { + config: Config, + main_channel: Arc, + response: Option, + iface: Option<(DeviceWriter, DeviceReader)>, +} + +impl SwitchUtil { + pub async fn new(config: Config) -> io::Result { + let main_channel = Arc::new(UdpSocket::bind("0.0.0.0:0").await?); + Ok(SwitchUtil { + config, + main_channel, + response: None, + iface: None, + }) + } + pub async fn connect(&mut self) -> Result { + match registration_handler::registration(&self.main_channel, self.config.server_address, + self.config.token.clone(), self.config.device_id.clone(), + self.config.name.clone()).await { + Ok(res) => { + let _ = self.response.insert(res.clone()); + Ok(res) + } + Err(e) => { + Err(e) + } + } + } + #[cfg(any(target_os = "android"))] + pub fn create_iface(&mut self, vpn_fd: i32) { + let (device_writer, device_reader) = tun_tap_device::create(vpn_fd); + let _ = self.iface.insert((device_writer, device_reader)); + } + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] + pub fn create_iface(&mut self) -> io::Result { + if self.iface.is_some() { + return Err(io::Error::from(io::ErrorKind::AlreadyExists)); + } + let response = match &self.response { + None => { + return Err(io::Error::from(io::ErrorKind::AlreadyExists)); + } + Some(res) => { + res + } + }; + let device_type = if self.config.tap { + #[cfg(windows)] + { + //删除switch的tun网卡避免ip冲突,因为非正常退出会保留网卡 + tun_tap_device::delete_device(tun_tap_device::DeviceType::Tun); + } + tun_tap_device::DeviceType::Tap + } else { + #[cfg(windows)] + { + //删除switch的tap网卡避免ip冲突,非正常退出会保留网卡 + tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap); + } + tun_tap_device::DeviceType::Tun + }; + let in_ips = self.config.in_ips.iter().map(|(dest, mask, _)| { (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)) }).collect::>(); + + let (device_writer, device_reader, driver_info) = tun_tap_device::create_device(device_type, response.virtual_ip, response.virtual_netmask, response.virtual_gateway, in_ips)?; + let _ = self.iface.insert((device_writer, device_reader)); + Ok(driver_info) + } + pub async fn build(self) -> crate::Result { + let response = match self.response { + None => { + return Err(Error::Stop("response None".to_string())); + } + Some(res) => { + res + } + }; + let (device_writer, device_reader) = match self.iface { + None => { + return Err(Error::Stop("iface None".to_string())); + } + Some(res) => { + res + } + }; + let config = self.config; + let switch_status_manager = SwitchStatusManger::new(); let cipher = if let Some(key) = &config.key { let key: &Key = key.into(); Some(Aes256Gcm::new(&key)) } else { None }; - let main_channel = Arc::new(UdpSocket::bind("0.0.0.0:0").await?); - let response = registration_handler::registration(&main_channel, config.server_address, config.token.clone(), config.device_id.clone(), config.name.clone()).await?; let (cone_sender, cone_receiver) = channel(3); let (symmetric_sender, symmetric_receiver) = channel(2); - let context = Context::new(main_channel, 1); + let context = Context::new(self.main_channel, 1); let punch = Punch::new(context.clone()); let idle = Idle::new(Duration::from_secs(16), context.clone()); let channel_sender = ChannelSender::new(context.clone()); - let register = Arc::new(registration_handler::Register::new(channel_sender.clone(), config.server_address, config.token.clone(), config.device_id.clone(), config.name.clone())); + let register = Arc::new(registration_handler::Register::new(channel_sender.clone(), + config.server_address, config.token.clone(), + config.device_id.clone(), config.name.clone())); let device_list: Arc)>> = Arc::new(Mutex::new((0, Vec::new()))); let peer_nat_info_map: Arc> = Arc::new(SkipMap::new()); let connect_status = Arc::new(AtomicCell::new(ConnectStatus::Connected)); - let virtual_ip = Ipv4Addr::from(response.virtual_ip); - let virtual_gateway = Ipv4Addr::from(response.virtual_gateway); - let virtual_netmask = Ipv4Addr::from(response.virtual_netmask); + let virtual_ip = response.virtual_ip; + let virtual_gateway = response.virtual_gateway; + let virtual_netmask = response.virtual_netmask; let local_ip = crate::nat::local_ip()?; let local_port = context.main_local_port()?; // NAT检测 - let nat_test = NatTest::new(config.nat_test_server.clone(), Ipv4Addr::from(response.public_ip), response.public_port as u16, local_ip, local_port); - let in_ips = config.in_ips.iter().map(|(dest, mask, _)| { (Ipv4Addr::from(*dest & *mask), Ipv4Addr::from(*mask)) }).collect::>(); + let nat_test = NatTest::new(config.nat_test_server.clone(), response.public_ip, response.public_port, local_ip, local_port); let out_ips = config.out_ips.iter().map(|(_, _, ip)| *ip).collect::>(); let out_external_route = ExternalRoute::new(config.out_ips); @@ -80,45 +174,30 @@ impl Switch { Some(ExternalRoute::new(config.in_ips)) }; let current_device = Arc::new(AtomicCell::new(CurrentDeviceInfo::new(virtual_ip, virtual_gateway, virtual_netmask, config.server_address))); - let ip_proxy_map = if out_ips.is_empty(){ - None - }else{ - Some(crate::ip_proxy::init_proxy(channel_sender.clone(), out_ips, current_device.clone()).await?) - }; - let (device_writer, igmp_server) = if config.tap { - #[cfg(windows)] - { - //删除switch的tun网卡避免ip冲突,因为非正常退出会保留网卡 - tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap); - } - let (tap_writer, tap_reader) = tun_tap_device::create_device(tun_tap_device::DeviceType::Tap, virtual_ip, virtual_netmask, virtual_gateway, in_ips)?; - let igmp_server = if config.simulate_multicast { - Some(IgmpServer::new(tap_writer.clone())) - } else { - None - }; - //tap数据处理 - tap_handler::start(channel_sender.clone(), tap_reader.clone(), tap_writer.clone(), - igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone()); - (tap_writer, igmp_server) + let (tcp_proxy, udp_proxy, ip_proxy_map) = if out_ips.is_empty() { + (None, None, None) } else { - #[cfg(windows)] - { - //删除switch的tap网卡避免ip冲突,非正常退出会保留网卡 - tun_tap_device::delete_device(tun_tap_device::DeviceType::Tap); - } - // tun通道 - let (tun_writer, tun_reader) = tun_tap_device::create_device(tun_tap_device::DeviceType::Tun, virtual_ip, virtual_netmask, virtual_gateway, in_ips)?; - let igmp_server = if config.simulate_multicast { - Some(IgmpServer::new(tun_writer.clone())) - } else { - None - }; - //tun数据接收处理 - tun_handler::start(channel_sender.clone(), tun_reader.clone(), tun_writer.clone(), - igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone()); - (tun_writer, igmp_server) + let (tcp_proxy, udp_proxy, ip_proxy_map) = crate::ip_proxy::init_proxy(channel_sender.clone(), out_ips, current_device.clone()).await?; + (Some(tcp_proxy), Some(udp_proxy), Some(ip_proxy_map)) }; + + let igmp_server = if config.simulate_multicast { + Some(IgmpServer::new(device_writer.clone())) + } else { + None + }; + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] + if config.tap { + tap_handler::start(switch_status_manager.worker(), channel_sender.clone(), device_reader, device_writer.clone(), + igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone()); + } else { + tun_handler::start(switch_status_manager.worker(), channel_sender.clone(), device_reader, device_writer.clone(), + igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone()); + } + #[cfg(any(target_os = "android"))] + tun_handler::start(switch_status_manager.worker(), channel_sender.clone(), device_reader, device_writer.clone(), + igmp_server.clone(), current_device.clone(), in_external_route, ip_proxy_map.clone(), cipher.clone()); + //外部数据接收处理 let channel_recv_handler = ChannelDataHandler::new(current_device.clone(), device_list.clone(), register.clone(), nat_test.clone(), igmp_server, @@ -126,27 +205,53 @@ impl Switch { peer_nat_info_map.clone(), ip_proxy_map, out_external_route, cone_sender, symmetric_sender, cipher); let channel = Channel::new(context.clone(), channel_recv_handler); + let channel_worker = switch_status_manager.worker(); + //数据接收 thread::spawn(move || { tokio::runtime::Builder::new_multi_thread() .enable_all() .build().unwrap() - .block_on(channel.start(14, 60)); + .block_on(async move { + if let Some(tcp_proxy) = tcp_proxy { + tokio::spawn(tcp_proxy.start()); + } + if let Some(udp_proxy) = udp_proxy { + tokio::spawn(udp_proxy.start()); + } + channel.start(channel_worker, 14, 65).await; + }); }); + { + let other_worker = switch_status_manager.worker(); + let nat_test = nat_test.clone(); + let device_list = device_list.clone(); + let current_device = current_device.clone(); + //其他任务处理 + thread::spawn(move || { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build().unwrap() + .block_on(async move { + // 定时心跳 + heartbeat_handler::start_heartbeat(other_worker.clone(), channel_sender.clone(), device_list.clone(), current_device.clone()); + // 空闲检查 + heartbeat_handler::start_idle(other_worker.clone(), idle, channel_sender.clone()); + // 打洞处理 + punch_handler::start(other_worker.clone(), cone_receiver, punch.clone(), current_device.clone()); + punch_handler::start(other_worker.clone(), symmetric_receiver, punch, current_device.clone()); + punch_handler::start_punch(other_worker.clone(), nat_test.clone(), + device_list.clone(), channel_sender.clone(), + current_device.clone()).await; + }); + }); + } context.switch(nat_test.nat_info().nat_type); - // 定时心跳 - heartbeat_handler::start_heartbeat(channel_sender.clone(), device_list.clone(), current_device.clone()).await; - // 空闲检查 - heartbeat_handler::start_idle(idle, channel_sender.clone()).await; - // 打洞处理 - punch_handler::start(cone_receiver, punch.clone(), current_device.clone()).await; - punch_handler::start(symmetric_receiver, punch, current_device.clone()).await; - punch_handler::start_punch(nat_test.clone(), device_list.clone(), channel_sender.clone(), current_device.clone()).await; - - log::info!("switch启动成功"); Ok(Switch { name: config.name, current_device, context, + switch_status_manager, + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] device_writer, nat_test, device_list, @@ -189,9 +294,15 @@ impl Switch { } pub fn stop(&self) -> io::Result<()> { self.context.close(); + self.switch_status_manager.stop_all(); + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] self.device_writer.close()?; Ok(()) } + pub async fn wait_stop(&mut self) { + self.switch_status_manager.wait().await; + let _ = self.stop(); + } } #[derive(Clone, Debug)] @@ -208,7 +319,6 @@ pub struct Config { pub simulate_multicast: bool, } -use sha2::Digest; impl Config { pub fn new(tap: bool, token: String, diff --git a/switch/src/core/status.rs b/switch/src/core/status.rs new file mode 100644 index 0000000..9333ed7 --- /dev/null +++ b/switch/src/core/status.rs @@ -0,0 +1,85 @@ +use std::sync::Arc; +use tokio::sync::watch; +use tokio::sync::watch::{Receiver, Sender}; +use crate::util::wait::WaitGroup; + +#[derive(Copy, Clone, Eq, PartialEq)] +pub enum SwitchStatus { + Starting, + Stopping, +} + +pub struct SwitchWorker { + wg: WaitGroup, + status_s: Arc>, + status_r: Receiver, +} + +impl Clone for SwitchWorker { + fn clone(&self) -> Self { + self.wg.add(); + SwitchWorker { + wg: self.wg.clone(), + status_s: self.status_s.clone(), + status_r: self.status_r.clone(), + } + } +} + +impl Drop for SwitchWorker { + fn drop(&mut self) { + self.wg.done(); + } +} + +impl SwitchWorker { + pub fn stop_all(&self) { + let _ = self.status_s.send(SwitchStatus::Stopping); + } + pub async fn stop_wait(&mut self) { + loop { + if *self.status_r.borrow() == SwitchStatus::Stopping { + return; + } + match self.status_r.changed().await { + Ok(_) => { + if *self.status_r.borrow() == SwitchStatus::Stopping { + return; + } + } + Err(_) => { return; } + } + } + } +} + +pub struct SwitchStatusManger { + wg: WaitGroup, + status_s: Arc>, + status_r: Receiver, +} + +impl SwitchStatusManger { + pub fn new() -> Self { + let (status_s, status_r) = watch::channel(SwitchStatus::Starting); + Self { + wg: WaitGroup::new(), + status_s: Arc::new(status_s), + status_r, + } + } + pub fn stop_all(&self) { + let _ = self.status_s.send(SwitchStatus::Stopping); + } + pub async fn wait(&mut self) { + self.wg.wait().await + } + pub fn worker(&self) -> SwitchWorker { + self.wg.add(); + SwitchWorker { + wg: self.wg.clone(), + status_s: self.status_s.clone(), + status_r: self.status_r.clone(), + } + } +} diff --git a/switch/src/core/sync.rs b/switch/src/core/sync.rs new file mode 100644 index 0000000..bcaf481 --- /dev/null +++ b/switch/src/core/sync.rs @@ -0,0 +1,64 @@ +use std::io; +use std::ops::Deref; +use std::time::Duration; +use tokio::runtime::Runtime; +use crate::core::{Config, Switch, SwitchUtil}; +use crate::handle::registration_handler::{RegResponse, ReqEnum}; + +pub struct SwitchUtilSync { + switch_util: SwitchUtil, + runtime: Runtime, +} + +pub struct SwitchSync { + switch: Switch, + runtime: Runtime, +} + +impl SwitchUtilSync { + pub fn new(config: Config) -> io::Result { + let runtime = tokio::runtime::Builder::new_current_thread().enable_all().build().unwrap(); + let switch_util = runtime.block_on(SwitchUtil::new(config))?; + Ok(SwitchUtilSync { + switch_util, + runtime, + }) + } + pub fn connect(&mut self) -> Result { + self.runtime.block_on(self.switch_util.connect()) + } + #[cfg(any(target_os = "android"))] + pub fn create_iface(&mut self, vpn_fd: i32) { + self.switch_util.create_iface(vpn_fd) + } + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] + pub fn create_iface(&mut self) -> io::Result { + self.switch_util.create_iface() + } + pub fn build(self) -> crate::Result { + let runtime = self.runtime; + let switch = runtime.block_on(self.switch_util.build())?; + Ok(SwitchSync { + switch, + runtime, + }) + } +} + +impl SwitchSync { + pub fn wait_stop(&mut self) { + self.runtime.block_on(self.switch.wait_stop()) + } + pub fn wait_stop_ms(&mut self, ms: u64) -> bool { + self.runtime.block_on(tokio::time::timeout(Duration::from_millis(ms), + self.switch.wait_stop())).is_ok() + } +} + +impl Deref for SwitchSync { + type Target = Switch; + + fn deref(&self) -> &Self::Target { + &self.switch + } +} \ No newline at end of file diff --git a/switch/src/handle/heartbeat_handler.rs b/switch/src/handle/heartbeat_handler.rs index 44a36ff..7518edf 100644 --- a/switch/src/handle/heartbeat_handler.rs +++ b/switch/src/handle/heartbeat_handler.rs @@ -9,20 +9,26 @@ use rand::prelude::SliceRandom; use crate::channel::idle::Idle; use crate::channel::Route; use crate::channel::sender::ChannelSender; +use crate::core::status::SwitchWorker; use crate::handle::{CurrentDeviceInfo, PeerDeviceInfo}; use crate::protocol::control_packet::PingPacket; use crate::protocol::{control_packet, NetPacket, Protocol, Version}; -pub async fn start_idle(idle: Idle, sender: ChannelSender) { +pub fn start_idle(mut worker: SwitchWorker, idle: Idle, sender: ChannelSender) { tokio::spawn(async move { - match start_idle_(idle, sender).await { - Ok(_) => {} - Err(e) => { - log::warn!("空闲检测任务停止:{:?}", e); + tokio::select! { + _=worker.stop_wait()=>{ + return; + } + rs=start_idle_(idle, sender)=>{ + if let Err(e) = rs { + log::warn!("空闲检测任务停止:{:?}", e); + } } } + worker.stop_all(); }); } @@ -38,15 +44,24 @@ async fn start_idle_(idle: Idle, sender: ChannelSender) -> io::Result<()> { } } -pub async fn start_heartbeat( +pub fn start_heartbeat( + mut worker: SwitchWorker, sender: ChannelSender, device_list: Arc)>>, current_device: Arc>, ) { tokio::spawn(async move { - if let Err(e) = start_heartbeat_(sender, device_list, current_device).await { - log::warn!("心跳任务停止:{:?}", e); + tokio::select! { + _=worker.stop_wait()=>{ + return; + } + rs=start_heartbeat_(sender, device_list, current_device)=>{ + if let Err(e) = rs { + log::warn!("心跳任务停止:{:?}", e); + } + } } + worker.stop_all(); }); } @@ -70,6 +85,9 @@ async fn start_heartbeat_( net_packet.first_set_ttl(2); let mut count = 0; loop { + if sender.is_close() { + return Ok(()); + } let current_device = current_device.load(); net_packet.set_source(current_device.virtual_ip()); { @@ -103,7 +121,7 @@ async fn start_heartbeat_( } } else { //没有直连路由则发送到网关 - let _ = sender.try_send_main(net_packet.buffer(), current_device.connect_server); + let _ = sender.send_main(net_packet.buffer(), current_device.connect_server).await; continue; } diff --git a/switch/src/handle/mod.rs b/switch/src/handle/mod.rs index 19bc083..d4359eb 100644 --- a/switch/src/handle/mod.rs +++ b/switch/src/handle/mod.rs @@ -38,7 +38,7 @@ impl PeerDeviceInfo { } } -#[derive(Copy, Clone, Debug, Eq, PartialEq)] +#[derive(Copy, Clone, Debug, Eq, PartialEq,Ord, PartialOrd)] pub enum PeerDeviceStatus { Online, Offline, diff --git a/switch/src/handle/punch_handler.rs b/switch/src/handle/punch_handler.rs index eec450f..360a598 100644 --- a/switch/src/handle/punch_handler.rs +++ b/switch/src/handle/punch_handler.rs @@ -13,10 +13,17 @@ use std::io; use tokio::sync::mpsc::Receiver; use crate::channel::punch::{NatInfo, Punch}; use crate::channel::sender::ChannelSender; +use crate::core::status::SwitchWorker; -pub async fn start(receiver: Receiver<(Ipv4Addr, NatInfo)>, punch: Punch, current_device: Arc>) { +pub fn start(mut worker: SwitchWorker, receiver: Receiver<(Ipv4Addr, NatInfo)>, punch: Punch, current_device: Arc>) { tokio::spawn(async move { - start0(receiver, punch, current_device).await; + tokio::select! { + _=start0(receiver, punch, current_device)=>{} + _=worker.stop_wait()=>{ + return; + } + } + worker.stop_all(); }); } @@ -47,54 +54,62 @@ async fn start_( } pub async fn start_punch( + mut worker: SwitchWorker, nat_test: NatTest, device_list: Arc)>>, sender: ChannelSender, current_device: Arc>, ) { - tokio::spawn(async move { - if let Err(e) = start_punch_(nat_test, device_list, sender, current_device).await { - log::warn!("打洞处理任务停止 {:?}", e); - } - }); -} - -async fn start_punch_( - nat_test: NatTest, - device_list: Arc)>>, - sender: ChannelSender, - current_device: Arc>, -) -> crate::Result<()> { let mut num = 0; let sleep_time = [3, 5, 7, 11, 13, 17, 19, 23, 29]; loop { if sender.is_close() { - return Ok(()); + break; } - 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.virtual_ip <= current_device.virtual_ip { - continue; + tokio::select! { + rs= start_punch_(Duration::from_secs(sleep_time[num % sleep_time.len()]),&nat_test, &device_list, &sender, ¤t_device)=>{ + if let Err(e) = rs { + log::warn!("打洞处理任务异常 {:?}", e); } - if !sender.need_punch(&info.virtual_ip) { - continue; - } - count += 1; - if count > 2 { - break; - } - let buf = punch_packet(current_device.virtual_ip(), &nat_info, info.virtual_ip)?; - let _ = sender.send_main(&buf, current_device.connect_server).await; + } + _=worker.stop_wait()=>{ + break; } } num += 1; - tokio::time::sleep(Duration::from_secs(sleep_time[num % sleep_time.len()])).await; } + + worker.stop_all(); +} + +async fn start_punch_( + sleep_time: Duration, + nat_test: &NatTest, + device_list: &Arc)>>, + sender: &ChannelSender, + current_device: &Arc>, +) -> crate::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.virtual_ip <= current_device.virtual_ip { + continue; + } + if !sender.need_punch(&info.virtual_ip) { + continue; + } + count += 1; + if count > 2 { + break; + } + let buf = punch_packet(current_device.virtual_ip(), &nat_info, info.virtual_ip)?; + let _ = sender.send_main(&buf, current_device.connect_server).await; + } + tokio::time::sleep(sleep_time).await; + Ok(()) } pub fn punch_packet( diff --git a/switch/src/handle/recv_handler.rs b/switch/src/handle/recv_handler.rs index 20b2d77..f0cc7e2 100644 --- a/switch/src/handle/recv_handler.rs +++ b/switch/src/handle/recv_handler.rs @@ -245,7 +245,7 @@ impl ChannelDataHandler { ipv4.set_destination_ip(destination); ipv4.update_checksum(); ip_proxy_map.tcp_proxy_map.insert(SocketAddrV4::new(source, source_port), - (SocketAddrV4::new(gate_way, 0), SocketAddrV4::new(dest_ip, dest_port))); + (SocketAddrV4::new(gate_way, 0), SocketAddrV4::new(dest_ip, dest_port))); } ipv4::protocol::Protocol::Udp => { let dest_ip = ipv4.destination_ip(); @@ -258,7 +258,7 @@ impl ChannelDataHandler { ipv4.set_destination_ip(destination); ipv4.update_checksum(); ip_proxy_map.udp_proxy_map.insert(SocketAddrV4::new(source, source_port), - (SocketAddrV4::new(gate_way, 0), SocketAddrV4::new(dest_ip, dest_port))); + (SocketAddrV4::new(gate_way, 0), SocketAddrV4::new(dest_ip, dest_port))); } ipv4::protocol::Protocol::Icmp => { let dest_ip = ipv4.destination_ip(); @@ -329,11 +329,14 @@ impl ChannelDataHandler { if current_ip != new_ip { // ip发生变化 log::info!("ip发生变化,old_ip:{:?},new_ip:{:?}",current_ip,new_ip); - let old_netmask = current_device.virtual_netmask; - let old_gateway = current_device.virtual_gateway(); + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] + let old_netmask = current_device.virtual_netmask; + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] + let old_gateway = current_device.virtual_gateway(); let virtual_ip = Ipv4Addr::from(response.virtual_ip); let virtual_gateway = Ipv4Addr::from(response.virtual_gateway); let virtual_netmask = Ipv4Addr::from(response.virtual_netmask); + #[cfg(any(target_os = "linux", target_os = "macos", target_os = "windows"))] self.device_writer.change_ip(virtual_ip, virtual_netmask, virtual_gateway, old_netmask, old_gateway)?; let new_current_device = CurrentDeviceInfo::new(virtual_ip, virtual_gateway, virtual_netmask, current_device.connect_server); @@ -346,7 +349,7 @@ impl ChannelDataHandler { service_packet::Protocol::PollDeviceList => {} service_packet::Protocol::PushDeviceList => { let device_list_t = DeviceList::parse_from_bytes(net_packet.payload())?; - let ip_list = device_list_t + let ip_list: Vec = device_list_t .device_info_list .into_iter() .map(|info| { @@ -357,6 +360,10 @@ impl ChannelDataHandler { ) }) .collect(); + let route = Route::from(*route_key, 2, 99); + for x in &ip_list { + context.add_route_if_absent(x.virtual_ip, route); + } let mut dev = self.device_list.lock(); if dev.0 != device_list_t.epoch as u16 { dev.0 = device_list_t.epoch as u16; @@ -434,7 +441,6 @@ impl ChannelDataHandler { } } ControlPacket::PunchRequest => { - // log::info!("PunchRequest route_key:{:?}",route_key); //回应 net_packet.set_transport_protocol(control_packet::Protocol::PunchResponse.into()); net_packet.set_source(current_device.virtual_ip()); diff --git a/switch/src/handle/registration_handler.rs b/switch/src/handle/registration_handler.rs index 8a796a1..778a0ed 100644 --- a/switch/src/handle/registration_handler.rs +++ b/switch/src/handle/registration_handler.rs @@ -1,5 +1,5 @@ use std::io; -use std::net::SocketAddr; +use std::net::{Ipv4Addr, SocketAddr}; use std::time::{Duration, Instant}; use crossbeam_utils::atomic::AtomicCell; @@ -7,11 +7,28 @@ use protobuf::Message; use tokio::net::UdpSocket; use crate::channel::sender::ChannelSender; -use crate::error::*; use crate::proto::message::{RegistrationRequest, RegistrationResponse}; use crate::protocol::error_packet::InErrorPacket; use crate::protocol::{service_packet, NetPacket, Protocol, Version, MAX_TTL}; +pub enum ReqEnum { + TokenError, + AddressExhausted, + Timeout, + ServerError(String), + Other(String), +} + +#[derive(Clone, Debug)] +pub struct RegResponse { + pub virtual_ip: Ipv4Addr, + pub virtual_gateway: Ipv4Addr, + pub virtual_netmask: Ipv4Addr, + pub epoch: u32, + pub public_ip: Ipv4Addr, + pub public_port: u16, +} + ///向中继服务器注册,token标识一个虚拟网关,device_id防止多次注册时得到的ip不一致 pub async fn registration( main_channel: &UdpSocket, @@ -19,77 +36,90 @@ pub async fn registration( token: String, device_id: String, name: String, -) -> Result { +) -> Result { let request_packet = - registration_request_packet(token.clone(), device_id.clone(), name.clone(), false)?; + registration_request_packet(token.clone(), device_id.clone(), name.clone(), false).unwrap(); let buf = request_packet.buffer(); let mut recv_buf = [0u8; 10240]; - let mut count = 0; - loop { - match main_channel.send_to(buf, server_address).await { - Ok(_) => { - match tokio::time::timeout(Duration::from_millis(300), main_channel.recv_from(&mut recv_buf)).await { - Ok(rs) => { - match rs { - Ok((len, addr)) => { - if server_address == addr { - let net_packet = NetPacket::new(&recv_buf[..len])?; - match net_packet.protocol() { - Protocol::Service => { - match service_packet::Protocol::from(net_packet.transport_protocol()) { - service_packet::Protocol::RegistrationResponse => { - let response = RegistrationResponse::parse_from_bytes(net_packet.payload())?; - return Ok(response); + return match main_channel.send_to(buf, server_address).await { + Ok(_) => { + match tokio::time::timeout(Duration::from_millis(300), main_channel.recv_from(&mut recv_buf)).await { + Ok(rs) => { + match rs { + Ok((len, addr)) => { + if server_address == addr { + let net_packet = match NetPacket::new(&recv_buf[..len]) { + Ok(net_packet) => { + net_packet + } + Err(e) => { + return Err(ReqEnum::ServerError(format!("{}",e))) + } + }; + match net_packet.protocol() { + Protocol::Service => { + match service_packet::Protocol::from(net_packet.transport_protocol()) { + service_packet::Protocol::RegistrationResponse => { + match RegistrationResponse::parse_from_bytes(net_packet.payload()) { + Ok(response) => { + Ok(RegResponse { + virtual_ip: Ipv4Addr::from(response.virtual_ip), + virtual_gateway: Ipv4Addr::from(response.virtual_gateway), + virtual_netmask: Ipv4Addr::from(response.virtual_netmask), + epoch: response.epoch, + public_ip: Ipv4Addr::from(response.public_ip), + public_port: response.public_port as u16, + }) + } + Err(_) => { + Err(ReqEnum::ServerError("invalid data".to_string())) + } } - _ => println!("响应数据错误"), + } + _ => { + Err(ReqEnum::ServerError("invalid data".to_string())) } } - Protocol::Error => { - match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload()) { - Ok(e) => match e { - InErrorPacket::TokenError => return Err(Error::Stop("token错误".to_string())), - InErrorPacket::Disconnect => { - println!("断开连接"); + } + Protocol::Error => { + match InErrorPacket::new(net_packet.transport_protocol(), net_packet.payload()) { + Ok(e) => match e { + InErrorPacket::TokenError => Err(ReqEnum::TokenError), + InErrorPacket::Disconnect => { + Err(ReqEnum::ServerError("disconnect".to_string())) + } + InErrorPacket::AddressExhausted => { + Err(ReqEnum::AddressExhausted) + } + InErrorPacket::OtherError(e) => match e.message() { + Ok(str) => { + Err(ReqEnum::ServerError(str)) } - InErrorPacket::AddressExhausted => { - println!("地址用尽"); - log::warn!("地址用尽"); - } - InErrorPacket::OtherError(e) => match e.message() { - Ok(str) => { - println!("其他异常:{:?}", str); - log::warn!("其他异常{:?}",str); - } - Err(e) => println!("其他异常:{:?}", e), - }, + Err(e) => Err(ReqEnum::Other(format!("{}", e))), }, - Err(e) => println!("数据解析异常:{:?}", e), - } + }, + Err(e) => Err(ReqEnum::Other(format!("{}", e))), } - _ => println!("响应数据错误"), - }; + } + _ => Err(ReqEnum::ServerError("invalid data".to_string())), } - } - Err(e) => { - println!("接收服务器数据失败:{:?}", e); - log::warn!("接收服务器数据失败:{:?}",e); + } else { + Err(ReqEnum::Other(format!("invalid data,from {}", addr))) } } - } - Err(_) => { - println!("接收超时"); - log::warn!("接收超时"); + Err(e) => { + Err(ReqEnum::Other(format!("receiver error:{}", e))) + } } } - } - Err(e) => { - println!("发送数据到服务器失败:{:?}", e); - log::warn!("发送数据到服务器失败:{:?}",e); + Err(_) => { + Err(ReqEnum::Timeout) + } } } - count += 1; - println!("重试中(retrying)..."); - std::thread::sleep(Duration::from_secs(count % 10 + 1)); + Err(e) => { + Err(ReqEnum::Other(format!("send error:{}", e))) + } }; } diff --git a/switch/src/handle/tun_tap/mod.rs b/switch/src/handle/tun_tap/mod.rs index 0e001dc..84402ef 100644 --- a/switch/src/handle/tun_tap/mod.rs +++ b/switch/src/handle/tun_tap/mod.rs @@ -18,41 +18,33 @@ use crate::protocol; use crate::protocol::ip_turn_packet::BroadcastPacketEnd; pub mod tun_handler; +#[cfg(any(target_os = "linux", target_os = "macos",target_os = "windows"))] pub mod tap_handler; async fn broadcast(sender: &ChannelSender, net_packet: &mut NetPacket<&mut [u8]>, data_len: usize, current_device: &CurrentDeviceInfo) -> Result<()> { let mut peer_ips = Vec::with_capacity(8); let vec = sender.route_table_one(); let mut relay_count = 0; - let mut last_peer = None; + const MAX_COUNT: usize = u8::MAX as usize; for (peer_ip, route) in vec { if peer_ip == current_device.virtual_gateway { continue; } - if peer_ips.len() < u8::MAX as usize && route.is_p2p() + if peer_ips.len() == MAX_COUNT { + break; + } + if route.is_p2p() && sender.send_by_key(&net_packet.buffer()[..data_len], &route.route_key()).await.is_ok() { peer_ips.push(peer_ip); } else { relay_count += 1; - if relay_count == 1 { - last_peer = Some((peer_ip, route)); - } - if relay_count > 1 && peer_ips.len() == u8::MAX as usize { - break; - } } } - if relay_count == 0 && !peer_ips.is_empty() { + if relay_count == 0 && !peer_ips.is_empty() && peer_ips.len() != MAX_COUNT { //不需要转发 return Ok(()); } - if relay_count == 1 && !net_packet.is_encrypt() { - //只有一个目标,并且没加密 - let (peer_ip, route) = last_peer.unwrap(); - net_packet.set_destination(peer_ip); - sender.send_by_key(&net_packet.buffer()[..data_len], &route.route_key()).await?; - return Ok(()); - } + if peer_ips.is_empty() { sender.send_main(&net_packet.buffer()[..data_len], current_device.connect_server).await?; } else { @@ -71,7 +63,7 @@ async fn multicast(igmp_server: &IgmpServer, multicast_addr: Ipv4Addr, sender: & let mut peer_ips = Vec::with_capacity(8); let vec = sender.route_table_one(); let mut relay_count = 0; - let mut last_peer = None; + const MAX_COUNT: usize = u8::MAX as usize; if let Some(members) = igmp_server.load(&multicast_addr) { for (peer_ip, route) in vec { if peer_ip == current_device.virtual_gateway { @@ -79,32 +71,22 @@ async fn multicast(igmp_server: &IgmpServer, multicast_addr: Ipv4Addr, sender: & } let is_send = { members.read().is_send(&peer_ip) }; if is_send { - if peer_ips.len() < u8::MAX as usize && route.is_p2p() + if peer_ips.len() == MAX_COUNT { + break; + } + if route.is_p2p() && sender.send_by_key(&net_packet.buffer()[..data_len], &route.route_key()).await.is_ok() { peer_ips.push(peer_ip); } else { relay_count += 1; - if relay_count == 1 { - last_peer = Some((peer_ip, route)); - } - if relay_count > 1 && peer_ips.len() == u8::MAX as usize { - break; - } } } } } - if relay_count == 0 && !peer_ips.is_empty() { + if relay_count == 0 && !peer_ips.is_empty() && peer_ips.len() != MAX_COUNT { //不需要转发 return Ok(()); } - if relay_count == 1 && !net_packet.is_encrypt() { - //只有一个目标,并且没加密 - let (peer_ip, route) = last_peer.unwrap(); - net_packet.set_destination(peer_ip); - sender.send_by_key(&net_packet.buffer()[..data_len], &route.route_key()).await?; - return Ok(()); - } if peer_ips.is_empty() { sender.send_main(&net_packet.buffer()[..data_len], current_device.connect_server).await?; } else { diff --git a/switch/src/handle/tun_tap/tap_handler.rs b/switch/src/handle/tun_tap/tap_handler.rs index fbe30d9..06d8f11 100644 --- a/switch/src/handle/tun_tap/tap_handler.rs +++ b/switch/src/handle/tun_tap/tap_handler.rs @@ -10,13 +10,14 @@ use packet::icmp::Kind; use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; use crate::channel::sender::ChannelSender; +use crate::core::status::SwitchWorker; use crate::external_route::ExternalRoute; use crate::handle::CurrentDeviceInfo; use crate::igmp_server::IgmpServer; use crate::ip_proxy::IpProxyMap; use crate::tun_tap_device::{DeviceReader, DeviceWriter}; -pub fn start(sender: ChannelSender, +pub fn start(worker: SwitchWorker, sender: ChannelSender, device_reader: DeviceReader, device_writer: DeviceWriter, igmp_server: Option, @@ -24,7 +25,7 @@ pub fn start(sender: ChannelSender, ip_route: Option, ip_proxy_map: Option, cipher: Option) { - thread::Builder::new().name("tap-handler".into()).spawn(move || { + thread::spawn(move || { tokio::runtime::Builder::new_current_thread() .enable_all().build().unwrap() .block_on(async move { @@ -33,8 +34,9 @@ pub fn start(sender: ChannelSender, current_device, ip_route, ip_proxy_map, cipher).await { log::warn!("tap:{:?}",e); } + worker.stop_all(); }); - }).unwrap(); + }); } async fn start_(sender: ChannelSender, diff --git a/switch/src/handle/tun_tap/tun_handler.rs b/switch/src/handle/tun_tap/tun_handler.rs index 4f09332..24618b1 100644 --- a/switch/src/handle/tun_tap/tun_handler.rs +++ b/switch/src/handle/tun_tap/tun_handler.rs @@ -9,6 +9,7 @@ use packet::icmp::icmp::IcmpPacket; use packet::ip::ipv4; use packet::ip::ipv4::packet::IpV4Packet; use crate::channel::sender::ChannelSender; +use crate::core::status::SwitchWorker; use crate::error::*; use crate::external_route::ExternalRoute; @@ -36,7 +37,7 @@ fn icmp(device_writer: &DeviceWriter, mut ipv4_packet: IpV4Packet<&mut [u8]>) -> /// 接收tun数据,并且转发到udp上 #[inline] async fn handle(sender: &ChannelSender, data: &mut [u8], len: usize, device_writer: &DeviceWriter, igmp_server: &Option, current_device: CurrentDeviceInfo, - ip_route: &Option, proxy_map: &Option,cipher: &Option) -> Result<()> { + ip_route: &Option, proxy_map: &Option, cipher: &Option) -> Result<()> { let ipv4_packet = if let Ok(ipv4_packet) = IpV4Packet::new(&mut data[12..len]) { ipv4_packet } else { @@ -50,10 +51,10 @@ async fn handle(sender: &ChannelSender, data: &mut [u8], len: usize, device_writ if src_ip == dest_ip { return icmp(&device_writer, ipv4_packet); } - return crate::handle::tun_tap::base_handle(sender, data, len, igmp_server, current_device, ip_route, proxy_map,cipher).await; + return crate::handle::tun_tap::base_handle(sender, data, len, igmp_server, current_device, ip_route, proxy_map, cipher).await; } -pub fn start(sender: ChannelSender, +pub fn start(worker: SwitchWorker, sender: ChannelSender, device_reader: DeviceReader, device_writer: DeviceWriter, igmp_server: Option, @@ -61,15 +62,16 @@ pub fn start(sender: ChannelSender, ip_route: Option, ip_proxy_map: Option, cipher: Option) { - thread::Builder::new().name("tun-handler".into()).spawn(move || { - tokio::runtime::Builder::new_multi_thread() + thread::spawn(move || { + tokio::runtime::Builder::new_current_thread() .enable_all().build().unwrap() .block_on(async move { - if let Err(e) = start_(sender, device_reader, device_writer, igmp_server, current_device, ip_route, ip_proxy_map,cipher).await { + if let Err(e) = start_(sender, device_reader, device_writer, igmp_server, current_device, ip_route, ip_proxy_map, cipher).await { log::warn!("tun:{:?}",e); } + worker.stop_all(); }) - }).unwrap(); + }); } async fn start_(sender: ChannelSender, @@ -80,24 +82,17 @@ async fn start_(sender: ChannelSender, ip_route: Option, ip_proxy_map: Option, cipher: Option) -> io::Result<()> { + let mut buf = [0; 4096]; loop { - let mut buf = [0; 4096]; - let sender = sender.clone(); - let device_writer = device_writer.clone(); - let igmp_server = igmp_server.clone(); - let ip_route = ip_route.clone(); - let ip_proxy_map = ip_proxy_map.clone(); - let cipher = cipher.clone(); + if sender.is_close() { + return Ok(()); + } let len = device_reader.read(&mut buf[12..])? + 12; - let current_device = current_device.load(); - tokio::spawn(async move { - match handle(&sender, &mut buf, len, &device_writer, &igmp_server, current_device, &ip_route, &ip_proxy_map,&cipher).await { - Ok(_) => {} - Err(e) => { - log::warn!("{:?}", e) - } + match handle(&sender, &mut buf, len, &device_writer, &igmp_server,current_device.load(), &ip_route, &ip_proxy_map, &cipher).await { + Ok(_) => {} + Err(e) => { + log::warn!("{:?}", e) } - }); - + } } } diff --git a/switch/src/ip_proxy/mod.rs b/switch/src/ip_proxy/mod.rs index ea7d5ba..d4a59a0 100644 --- a/switch/src/ip_proxy/mod.rs +++ b/switch/src/ip_proxy/mod.rs @@ -45,30 +45,17 @@ impl IpProxyMap { } } -pub async fn init_proxy(sender: ChannelSender, bind_ips: Vec, current_device: Arc>) -> io::Result { +pub async fn init_proxy(sender: ChannelSender, bind_ips: Vec, current_device: Arc>) -> io::Result<(TcpProxy, UdpProxy, IpProxyMap)> { let mut icmp_sockets = HashMap::new(); let tcp_proxy_map: Arc> = Arc::new(SkipMap::new()); let udp_proxy_map: Arc> = Arc::new(SkipMap::new()); let icmp_proxy_map: Arc> = Arc::new(SkipMap::new()); - let (tcp_proxy_port, udp_proxy_port) = if !bind_ips.is_empty() { - let tcp_listener = TcpListener::bind("0.0.0.0:0").await?; - let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; - let tcp_proxy_port = tcp_listener.local_addr()?.port(); - let udp_proxy_port = udp_socket.local_addr()?.port(); - let tcp_proxy_map = tcp_proxy_map.clone(); - tokio::spawn(async { - let tcp_proxy = TcpProxy::new(tcp_listener, tcp_proxy_map); - tcp_proxy.start().await - }); - let udp_proxy_map = udp_proxy_map.clone(); - tokio::spawn(async { - let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map); - udp_proxy.start().await - }); - (tcp_proxy_port, udp_proxy_port) - } else { - (0, 0) - }; + let tcp_listener = TcpListener::bind("0.0.0.0:0").await?; + let udp_socket = UdpSocket::bind("0.0.0.0:0").await?; + let tcp_proxy_port = tcp_listener.local_addr()?.port(); + let udp_proxy_port = udp_socket.local_addr()?.port(); + let tcp_proxy = TcpProxy::new(tcp_listener, tcp_proxy_map.clone()); + let udp_proxy = UdpProxy::new(udp_socket, udp_proxy_map.clone()); for ip in bind_ips { let addr = SocketAddrV4::new(ip, 0); let icmp_proxy_map = icmp_proxy_map.clone(); @@ -79,12 +66,12 @@ pub async fn init_proxy(sender: ChannelSender, bind_ips: Vec, current_ }); } - Ok(IpProxyMap { + Ok((tcp_proxy, udp_proxy, IpProxyMap { tcp_proxy_port, udp_proxy_port, tcp_proxy_map, udp_proxy_map, icmp_proxy_map, icmp_sockets, - }) + })) } \ No newline at end of file diff --git a/switch/src/lib.rs b/switch/src/lib.rs index eebd814..b4a1401 100644 --- a/switch/src/lib.rs +++ b/switch/src/lib.rs @@ -13,3 +13,4 @@ pub mod igmp_server; pub mod tun_tap_device; pub mod core; pub mod channel; +pub mod util; diff --git a/switch/src/nat/check.rs b/switch/src/nat/check.rs index 5d53336..155a362 100644 --- a/switch/src/nat/check.rs +++ b/switch/src/nat/check.rs @@ -58,7 +58,6 @@ pub fn public_ip_list_( udp: &UdpSocket, addrs: &Vec, ) -> io::Result<(HashSet, u16, u16)> { - // println!("local port {:?}", udp.local_addr().unwrap().port()); udp.set_read_timeout(Some(Duration::from_millis(300)))?; let mut buf = [0u8; 128]; for addr in addrs { diff --git a/switch/src/tun_tap_device/android.rs b/switch/src/tun_tap_device/android.rs new file mode 100644 index 0000000..508877b --- /dev/null +++ b/switch/src/tun_tap_device/android.rs @@ -0,0 +1,42 @@ +use std::io; +use std::os::unix::io::RawFd; + +#[derive(Clone)] +pub struct DeviceWriter(RawFd); + +pub struct DeviceReader(RawFd); + +impl DeviceWriter { + pub fn write_ipv4_tun(&self, buf: &[u8]) -> io::Result<()> { + unsafe { + let amount = libc::write(self.0, buf.as_ptr() as *const _, buf.len() ); + if amount < 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) + } + } + ///写入ipv4数据,为了兼容其他代码,头部空了14个字节 + pub fn write_ipv4(&self, buf: &[u8]) -> io::Result<()> { + let buf = &buf[14..]; + self.write_ipv4_tun(buf) + } +} + +impl DeviceReader { + pub fn read(&self, buf: &mut [u8]) -> io::Result { + unsafe { + let amount = libc::read(self.0, buf.as_mut_ptr() as *mut _, buf.len() ); + + if amount < 0 { + return Err(io::Error::last_os_error()); + } + + Ok(amount as usize) + } + } +} + +pub fn create(fd: i32) -> (DeviceWriter, DeviceReader) { + (DeviceWriter(fd as _), DeviceReader(fd as _)) +} \ No newline at end of file diff --git a/switch/src/tun_tap_device/linux.rs b/switch/src/tun_tap_device/linux.rs index a317f1e..0055129 100644 --- a/switch/src/tun_tap_device/linux.rs +++ b/switch/src/tun_tap_device/linux.rs @@ -1,11 +1,11 @@ use std::io; use std::net::Ipv4Addr; -use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter}; +use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo}; use tun::Device; use parking_lot::Mutex; use std::process::Command; use std::sync::Arc; -use crate::tun_tap_device::unix::DeviceW; +use crate::tun_tap_device::linux_mac::DeviceW; impl DeviceWriter { pub fn change_ip(&self, address: Ipv4Addr, netmask: Ipv4Addr, @@ -56,8 +56,7 @@ pub fn create_device(device_type: DeviceType, netmask: Ipv4Addr, gateway: Ipv4Addr, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, -) -> io::Result<(DeviceWriter, DeviceReader)> { - println!("========网卡配置========"); +) -> io::Result<(DeviceWriter, DeviceReader,DriverInfo)> { let mut config = tun::Configuration::default(); config @@ -79,7 +78,6 @@ pub fn create_device(device_type: DeviceType, let reader = queue.reader(); let writer = queue.writer(); let name = dev.name(); - println!("name:{:?}", name); for (address, netmask) in &in_ips { add_route(name, *address, *netmask)?; } @@ -111,10 +109,16 @@ pub fn create_device(device_type: DeviceType, DeviceW::Tap((writer, mac)) } }; - println!("========TUN网卡配置========"); + let driver_info = DriverInfo { + device_type, + name:name.to_string(), + version:String::new(), + mac: None, + }; Ok(( DeviceWriter::new(device_w, Arc::new(Mutex::new(dev)), in_ips, address, packet_information), DeviceReader::new(reader), + driver_info, )) } diff --git a/switch/src/tun_tap_device/unix.rs b/switch/src/tun_tap_device/linux_mac.rs similarity index 97% rename from switch/src/tun_tap_device/unix.rs rename to switch/src/tun_tap_device/linux_mac.rs index b36d7e7..f1f035e 100644 --- a/switch/src/tun_tap_device/unix.rs +++ b/switch/src/tun_tap_device/linux_mac.rs @@ -6,9 +6,9 @@ use tun::platform::posix::{Reader, Writer}; use std::net::Ipv4Addr; use std::os::unix::io::AsRawFd; use crossbeam_utils::atomic::AtomicCell; -#[cfg(any(target_os = "linux", target_os = "android"))] +#[cfg(any(target_os = "linux"))] use tun::platform::linux::Device; -#[cfg(any(target_os = "macos", target_os = "ios"))] +#[cfg(any(target_os = "macos"))] use tun::platform::macos::Device; use parking_lot::Mutex; use packet::ethernet; @@ -128,7 +128,6 @@ impl DeviceWriter { } } -#[derive(Clone)] pub struct DeviceReader(Reader); impl DeviceReader { diff --git a/switch/src/tun_tap_device/mac.rs b/switch/src/tun_tap_device/mac.rs index b96e924..8146a0f 100644 --- a/switch/src/tun_tap_device/mac.rs +++ b/switch/src/tun_tap_device/mac.rs @@ -1,11 +1,11 @@ use std::io; use std::net::Ipv4Addr; -use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter}; +use crate::tun_tap_device::{DeviceReader, DeviceType, DeviceWriter, DriverInfo}; use tun::Device; use parking_lot::Mutex; use std::process::Command; use std::sync::Arc; -use crate::tun_tap_device::unix::DeviceW; +use crate::tun_tap_device::linux_mac::DeviceW; impl DeviceWriter { pub fn change_ip(&self, address: Ipv4Addr, netmask: Ipv4Addr, @@ -42,14 +42,13 @@ pub fn create_device(device_type: DeviceType, netmask: Ipv4Addr, gateway: Ipv4Addr, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, -) -> io::Result<(DeviceWriter, DeviceReader)> { +) -> io::Result<(DeviceWriter, DeviceReader,DriverInfo)> { match device_type { DeviceType::Tun => {} DeviceType::Tap => { unimplemented!() } } - println!("========TUN网卡配置========"); let mut config = tun::Configuration::default(); config @@ -74,11 +73,16 @@ pub fn create_device(device_type: DeviceType, let queue = dev.queue(0).unwrap(); let reader = queue.reader(); let writer = queue.writer(); - println!("name:{:?}", name); - println!("========TUN网卡配置========"); + let driver_info = DriverInfo { + device_type, + name:name.to_string(), + version:String::new(), + mac: None, + }; Ok(( DeviceWriter::new(DeviceW::Tun(writer), Arc::new(Mutex::new(dev)), in_ips, address, packet_information), DeviceReader::new(reader), + driver_info )) } diff --git a/switch/src/tun_tap_device/mod.rs b/switch/src/tun_tap_device/mod.rs index fabc12d..374b263 100644 --- a/switch/src/tun_tap_device/mod.rs +++ b/switch/src/tun_tap_device/mod.rs @@ -1,19 +1,25 @@ #[cfg(target_os = "windows")] -pub mod windows; -#[cfg(any(target_os = "linux", target_os = "android"))] -pub mod linux; +mod windows; +#[cfg(any(target_os = "linux"))] +mod linux; #[cfg(target_os = "macos")] -pub mod mac; -#[cfg(any(unix))] -pub mod unix; +mod mac; +#[cfg(any(target_os = "linux", target_os = "macos"))] +mod linux_mac; +#[cfg(target_os = "android")] +mod android; -#[cfg(any(target_os = "linux", target_os = "android"))] + +#[cfg(any(target_os = "linux"))] pub use linux::create_device; -#[cfg(any(target_os = "linux", target_os = "android"))] +#[cfg(any(target_os = "linux"))] pub use linux::delete_device; -#[cfg(any(unix))] -pub use unix::{DeviceWriter, DeviceReader}; - +#[cfg(target_os = "android")] +pub use android::create; +#[cfg(any(target_os = "linux", target_os = "macos"))] +pub use linux_mac::{DeviceWriter, DeviceReader}; +#[cfg(target_os = "android")] +pub use android::{DeviceWriter, DeviceReader}; #[cfg(target_os = "macos")] pub use mac::create_device; #[cfg(target_os = "macos")] @@ -26,7 +32,21 @@ pub use windows::delete_device; #[cfg(target_os = "windows")] pub use windows::{DeviceWriter, DeviceReader}; +#[derive(Copy, Clone, Debug, Eq, PartialEq)] pub enum DeviceType { Tun, Tap, +} + +impl DeviceType { + pub fn is_tun(&self) -> bool { + *self == DeviceType::Tun + } +} + +pub struct DriverInfo { + pub device_type: DeviceType, + pub name: String, + pub version: String, + pub mac: Option, } \ No newline at end of file diff --git a/switch/src/tun_tap_device/windows.rs b/switch/src/tun_tap_device/windows.rs index d9021d6..5c8068a 100644 --- a/switch/src/tun_tap_device/windows.rs +++ b/switch/src/tun_tap_device/windows.rs @@ -8,7 +8,7 @@ use parking_lot::Mutex; use packet::ethernet; use packet::ethernet::packet::EthernetPacket; use win_tun_tap::{IFace, TapDevice, TunDevice}; -use crate::tun_tap_device::DeviceType; +use crate::tun_tap_device::{DriverInfo, DeviceType}; pub const TUN_INTERFACE_NAME: &str = "Switch-Tun-V1"; pub const TUN_POOL_NAME: &str = "Switch-Tun-V1"; @@ -161,7 +161,6 @@ fn dest(ip: Ipv4Addr, mask: Ipv4Addr) -> Ipv4Addr { ]) } -#[derive(Clone)] pub struct DeviceReader { device: Arc, } @@ -199,9 +198,8 @@ fn create_tun( netmask: Ipv4Addr, gateway: Ipv4Addr, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, -) -> io::Result<(DeviceWriter, DeviceReader)> { +) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> { unsafe { - println!("========TUN网卡配置========"); match Library::new("wintun.dll") { Ok(lib) => match TunDevice::delete_for_name(lib, TUN_INTERFACE_NAME) { Ok(_) => { @@ -210,7 +208,6 @@ fn create_tun( Err(_) => {} }, Err(e) => { - log::error!("wintun.dll not found"); return Err(io::Error::new( io::ErrorKind::Other, format!("wintun.dll not found {:?}", e), @@ -240,8 +237,8 @@ fn create_tun( } } }; - println!("name:{:?}", tun_device.get_name()?); - println!("version:{:?}", tun_device.version()?); + let name = tun_device.get_name()?; + let version = format!("{:?}", tun_device.version()?); tun_device.set_ip(address, netmask)?; tun_device.set_metric(1)?; tun_device.set_mtu(1420)?; @@ -256,14 +253,21 @@ fn create_tun( tun_device.add_route(Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]), gateway, 1)?; delete_cache(); let device = Arc::new(Device::Tun(tun_device)); - println!("========TUN网卡配置========"); + let driver_info = DriverInfo { + device_type: DeviceType::Tun, + name, + version, + mac: None, + }; Ok(( DeviceWriter::new(device.clone(), in_ips, address), DeviceReader::new(device), + driver_info )) } } -fn delete_cache(){ + +fn delete_cache() { //清除路由缓存 let delete_cache = "netsh interface ip delete destinationcache"; let out = std::process::Command::new("cmd") @@ -271,7 +275,7 @@ fn delete_cache(){ .arg(delete_cache) .output() .unwrap(); - if !out.status.success(){ + if !out.status.success() { log::warn!("删除缓存失败:{:?}",out); } } @@ -293,8 +297,7 @@ fn create_tap( netmask: Ipv4Addr, gateway: Ipv4Addr, in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, -) -> io::Result<(DeviceWriter, DeviceReader)> { - println!("========TAP网卡配置========"); +) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> { let tap_device = match TapDevice::open(TAP_INTERFACE_NAME) { Ok(tap_device) => tap_device, Err(e) => { @@ -305,9 +308,9 @@ fn create_tap( } }; let mac = tap_device.get_mac()?; - println!("name:{:?}", tap_device.get_name()?); - println!("version:{:x?}", tap_device.get_version()?); - println!("mac:{:x?}", mac); + let name = tap_device.get_name()?; + let version = format!("{:?}", tap_device.get_version()?); + let mac_str = format!("mac:{:x?}", mac); tap_device.set_ip(address, netmask)?; tap_device.set_metric(1)?; tap_device.set_mtu(1420)?; @@ -321,10 +324,16 @@ fn create_tap( tap_device.add_route(Ipv4Addr::from([224, 0, 0, 0]), Ipv4Addr::from([240, 0, 0, 0]), gateway, 1)?; delete_cache(); let tap = Arc::new(Device::Tap((tap_device, mac))); - println!("========TAP网卡配置========"); + let driver_info = DriverInfo { + device_type: DeviceType::Tap, + name, + version, + mac: Some(mac_str), + }; Ok(( DeviceWriter::new(tap.clone(), in_ips, address), - DeviceReader::new(tap) + DeviceReader::new(tap), + driver_info )) } @@ -341,7 +350,7 @@ fn delete_tap() { pub fn create_device(device_type: DeviceType, address: Ipv4Addr, netmask: Ipv4Addr, gateway: Ipv4Addr, - in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, ) -> io::Result<(DeviceWriter, DeviceReader)> { + in_ips: Vec<(Ipv4Addr, Ipv4Addr)>, ) -> io::Result<(DeviceWriter, DeviceReader, DriverInfo)> { match device_type { DeviceType::Tun => { create_tun(address, netmask, gateway, in_ips) diff --git a/switch/src/util/mod.rs b/switch/src/util/mod.rs new file mode 100644 index 0000000..3f25590 --- /dev/null +++ b/switch/src/util/mod.rs @@ -0,0 +1 @@ +pub mod wait; \ No newline at end of file diff --git a/switch/src/util/wait.rs b/switch/src/util/wait.rs new file mode 100644 index 0000000..669f6e1 --- /dev/null +++ b/switch/src/util/wait.rs @@ -0,0 +1,44 @@ +use std::sync::Arc; +use std::sync::atomic::{AtomicIsize, Ordering}; +use tokio::sync::watch::{channel, Receiver, Sender}; + +#[derive(Clone)] +pub struct WaitGroup { + count: Arc, + receiver: Receiver, + sender: Arc>, +} + +impl WaitGroup { + pub fn new() -> Self { + let (sender, receiver) = channel(1); + Self { + count: Arc::new(Default::default()), + receiver, + sender: Arc::new(sender), + } + } + pub fn add(&self) { + let _ = self.count.fetch_add(1, Ordering::Relaxed); + } + pub fn done(&self) { + let i = self.count.fetch_sub(1, Ordering::Relaxed); + if i == 1 { + let _ = self.sender.send(0); + } + } + pub async fn wait(&mut self) { + loop { + if 0 == *self.receiver.borrow() { + return; + } + if self.receiver.changed().await.is_ok() { + if 0 == *self.receiver.borrow() { + return; + } + } else { + return; + } + } + } +} \ No newline at end of file