From b2b58a23c1d9a449457e038970074a8200acb512 Mon Sep 17 00:00:00 2001 From: lubeilin <1791778603@qq.com> Date: Sun, 10 Mar 2024 14:11:06 +0800 Subject: [PATCH] =?UTF-8?q?[mio]=20=E4=B8=BB=E9=80=9A=E9=81=93=E6=94=B9?= =?UTF-8?q?=E4=B8=BA=E5=BC=82=E6=AD=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- vnt/src/channel/context.rs | 111 ++++++++++++++---------- vnt/src/channel/mod.rs | 20 +++-- vnt/src/channel/udp_channel.rs | 153 ++++++++++++++++++++++++--------- 3 files changed, 194 insertions(+), 90 deletions(-) diff --git a/vnt/src/channel/context.rs b/vnt/src/channel/context.rs index eb2c60a..99b2d53 100644 --- a/vnt/src/channel/context.rs +++ b/vnt/src/channel/context.rs @@ -1,10 +1,10 @@ use std::collections::HashMap; -use std::io; use std::net::{Ipv4Addr, SocketAddr, SocketAddrV6, UdpSocket}; use std::ops::Deref; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::{Duration, Instant}; +use std::{io, thread}; use crossbeam_utils::atomic::AtomicCell; use parking_lot::RwLock; @@ -51,6 +51,7 @@ impl Context { state: AtomicBool::new(true), packet_loss_rate, packet_delay, + main_index: AtomicUsize::new(0), }; Self { inner: Arc::new(inner), @@ -89,6 +90,7 @@ pub struct ContextInner { packet_loss_rate: u32, //控制延迟 packet_delay: u32, + main_index: AtomicUsize, } impl ContextInner { @@ -222,9 +224,13 @@ impl ContextInner { //服务端地址只在重连时检测变化 self.send_tcp(buf, addr) } else { - self.send_main_udp(0, buf, addr) + self.send_main_udp(self.main_index.load(Ordering::Relaxed), buf, addr) } } + pub fn change_main_index(&self) { + let index = (self.main_index.load(Ordering::Relaxed) + 1) % self.main_udp_socket.len(); + self.main_index.store(index, Ordering::Relaxed); + } /// 此方法仅用于对称网络打洞 pub fn try_send_all(&self, buf: &[u8], addr: SocketAddr) { self.try_send_all_main(buf, addr); @@ -232,6 +238,7 @@ impl ContextInner { if let Err(e) = udp.send_to(buf, addr) { log::warn!("{:?},add={:?}", e, addr); } + thread::sleep(Duration::from_millis(1)); } } pub fn try_send_all_main(&self, buf: &[u8], mut addr: SocketAddr) { @@ -262,18 +269,39 @@ impl ContextInner { } } if self.packet_delay > 0 { - std::thread::sleep(Duration::from_millis(self.packet_delay as _)); + thread::sleep(Duration::from_millis(self.packet_delay as _)); } - if self.send_by_id(buf, id).is_err() && !self.route_table.use_channel_type.is_only_p2p() { - self.send_default(buf, server_addr) - } else { - Ok(()) + //优先发到直连到地址 + if let Err(e) = self.send_by_id(buf, id) { + if e.kind() != io::ErrorKind::NotFound { + log::warn!("{}:{:?}", id, e); + } + if !self.route_table.use_channel_type.is_only_p2p() { + //符合条件再发到服务器转发 + self.send_default(buf, server_addr)?; + } } + Ok(()) } /// 将数据发到指定id pub fn send_by_id(&self, buf: &[u8], id: &Ipv4Addr) -> io::Result<()> { - let route = self.route_table.get_route_by_id(id)?; - self.send_by_key(buf, route.route_key()) + let mut c = 0; + loop { + let route = self.route_table.get_route_by_id(c, id)?; + return if let Err(e) = self.send_by_key(buf, route.route_key()) { + //降低发送速率 + if e.kind() == io::ErrorKind::WouldBlock { + c += 1; + if c < 10 { + thread::sleep(Duration::from_micros(200)); + continue; + } + } + Err(e) + } else { + Ok(()) + }; + } } /// 将数据发到指定路由 pub fn send_by_key(&self, buf: &[u8], route_key: RouteKey) -> io::Result<()> { @@ -329,30 +357,17 @@ impl RouteTable { } impl RouteTable { - fn get_route_by_id(&self, id: &Ipv4Addr) -> io::Result { + fn get_route_by_id(&self, index: usize, id: &Ipv4Addr) -> io::Result { if let Some((_count, v)) = self.route_table.read().get(id) { - let len = v.len(); - if len == 0 { - return Err(io::Error::new(io::ErrorKind::NotFound, "route not found")); - } - // 因为列表是按延迟排序的,会一直变,直接取第一条是合理的 - let (route, time) = &v[0]; - - // 刚加入的或者长时间没通信的不使用 - if route.rt != DEFAULT_RT && time.load().elapsed() < Duration::from_secs(5) { - return Ok(*route); - } - // 如果指定路由不符合,则遍历路由表找到符合条件的 - if len > 1 { - for (route, time) in v[1..].iter() { - if route.rt != DEFAULT_RT && time.load().elapsed() < Duration::from_secs(5) { - return Ok(*route); - } + if self.first_latency { + if let Some((route, _)) = v.first() { + return Ok(*route); + } + } else { + let len = v.len(); + if len != 0 { + return Ok(v[index % len].0); } - } - //加一条保底 - if route.is_p2p() && route.rt != DEFAULT_RT { - return Ok(*route); } } Err(io::Error::new(io::ErrorKind::NotFound, "route not found")) @@ -367,7 +382,7 @@ impl RouteTable { // 限制通道类型 match self.use_channel_type { UseChannelType::P2p => { - if route.metric != 1 { + if !route.is_p2p() { return; } } @@ -379,7 +394,6 @@ impl RouteTable { .entry(id) .or_insert_with(|| (AtomicUsize::new(0), Vec::with_capacity(4))); let mut exist = false; - let mut p2p_num = 0; for (x, time) in list.iter_mut() { if x.metric < route.metric && !self.first_latency { //非优先延迟的情况下 不能比当前的路径更长 @@ -395,26 +409,25 @@ impl RouteTable { time.store(Instant::now()); break; } - if x.is_p2p() { - p2p_num += 1; - } } if exist { list.sort_by_key(|(k, _)| k.rt); - } else { - let limit_len = if self.first_latency { - self.channel_num - } else { - if p2p_num >= self.channel_num { - // p2p通道满员了则不再添加 + //如果延迟都稳定了,则去除多余通道 + for (route, _) in list.iter() { + if route.rt == DEFAULT_RT { return; } - if route.metric == 1 { + } + list.truncate(self.channel_num); + } else { + if !self.first_latency { + if route.is_p2p() { //非优先延迟的情况下 添加了直连的则排除非直连的 list.retain(|(k, _)| k.is_p2p()); } - self.channel_num - 1 }; + //增加路由表容量,避免波动 + let limit_len = self.channel_num * 2; list.sort_by_key(|(k, _)| k.rt); if list.len() > limit_len { list.truncate(limit_len); @@ -436,6 +449,16 @@ impl RouteTable { None } } + pub fn route_one_p2p(&self, id: &Ipv4Addr) -> Option { + if let Some((_, v)) = self.route_table.read().get(id) { + for (i, _) in v { + if i.is_p2p() { + return Some(*i); + } + } + } + None + } pub fn route_to_id(&self, route_key: &RouteKey) -> Option { let table = self.route_table.read(); for (k, (_, v)) in table.iter() { diff --git a/vnt/src/channel/mod.rs b/vnt/src/channel/mod.rs index 80fce49..6b0dd28 100644 --- a/vnt/src/channel/mod.rs +++ b/vnt/src/channel/mod.rs @@ -1,7 +1,6 @@ use std::io; use std::net::{SocketAddr, UdpSocket}; use std::str::FromStr; -use std::time::Duration; use crate::channel::context::Context; use crate::channel::handler::RecvChannelHandler; @@ -156,17 +155,26 @@ pub fn init_context( assert!(!ports.is_empty(), "not channel"); let mut udps = Vec::with_capacity(ports.len()); for port in &ports { - //监听v6+v4双栈,主通道使用同步io + //监听v6+v4双栈 let address: SocketAddr = format!("[::]:{}", port).parse().unwrap(); let socket = socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None)?; io_convert(socket.set_only_v6(false), |_| { format!("set_only_v6 failed: {}", &address) })?; + io_convert(socket.set_reuse_address(true), |_| { + format!("set_reuse_address failed: {}", &address) + })?; + io_convert(socket.set_send_buffer_size(2 * 1024 * 1024), |_| { + format!("set_send_buffer_size failed: {}", &address) + })?; + io_convert(socket.set_recv_buffer_size(2 * 1024 * 1024), |_| { + format!("set_recv_buffer_size failed: {}", &address) + })?; io_convert(socket.bind(&address.into()), |_| { format!("bind failed: {}", &address) })?; let main_channel: UdpSocket = socket.into(); - main_channel.set_write_timeout(Some(Duration::from_secs(5)))?; + main_channel.set_nonblocking(true)?; udps.push(main_channel); } let context = Context::new( @@ -185,7 +193,9 @@ pub fn init_context( io_convert(socket.set_only_v6(false), |_| { format!("set_only_v6 failed: {}", &address) })?; - + io_convert(socket.set_reuse_address(true), |_| { + format!("set_reuse_address failed: {}", &address) + })?; if let Err(e) = socket.bind(&address.into()) { if ports[0] == 0 { //端口可能冲突,则使用任意端口 @@ -199,7 +209,7 @@ pub fn init_context( io_convert(Err(e), |_| format!("bind failed: {}", &address))?; } } - socket.listen(2)?; + socket.listen(128)?; socket.set_nonblocking(true)?; socket.set_nodelay(false)?; let tcp_listener = mio::net::TcpListener::from_std(socket.into()); diff --git a/vnt/src/channel/udp_channel.rs b/vnt/src/channel/udp_channel.rs index 3957381..4b3843d 100644 --- a/vnt/src/channel/udp_channel.rs +++ b/vnt/src/channel/udp_channel.rs @@ -1,7 +1,6 @@ use std::collections::HashMap; -use std::net::UdpSocket as StdUdpSocket; -use std::net::{Ipv4Addr, SocketAddr}; use std::sync::mpsc::{sync_channel, Receiver}; +use std::sync::Arc; use std::{io, thread}; use mio::event::Source; @@ -23,15 +22,7 @@ pub fn udp_listen( where H: RecvChannelHandler, { - //根据通道数创建对应线程进行读取 - for index in 0..context.channel_num() { - main_udp_listen( - index, - stop_manager.clone(), - recv_handler.clone(), - context.clone(), - )?; - } + main_udp_listen(stop_manager.clone(), recv_handler.clone(), context.clone())?; sub_udp_listen(stop_manager, recv_handler, context) } @@ -146,7 +137,6 @@ where /// 阻塞监听 fn main_udp_listen( - index: usize, stop_manager: StopManager, recv_handler: H, context: Context, @@ -154,54 +144,135 @@ fn main_udp_listen( where H: RecvChannelHandler, { - let port = context.main_udp_socket[index].local_addr()?.port(); - let context_ = context.clone(); - let worker = stop_manager.add_listener(format!("main_udp_listen-{}", index), move || { - context_.stop(); - match StdUdpSocket::bind("127.0.0.1:0") { - Ok(udp) => { - if let Err(e) = udp.send_to( - b"stop", - SocketAddr::V4(std::net::SocketAddrV4::new(Ipv4Addr::LOCALHOST, port)), - ) { - log::error!("发送停止消息到udp失败:{:?}", e); - } - } - Err(e) => { - log::error!("发送停止-绑定udp失败:{:?}", e); - } + let poll = Poll::new()?; + let waker = Arc::new(Waker::new(poll.registry(), NOTIFY)?); + let _waker = waker.clone(); + let worker = stop_manager.add_listener("main_udp".into(), move || { + if let Err(e) = waker.wake() { + log::error!("{:?}", e); } })?; thread::Builder::new() - .name("main_udp读事件处理线程".into()) + .name("main_udp".into()) .spawn(move || { - if let Err(e) = main_udp_listen0(index, recv_handler, context) { + if let Err(e) = main_udp_listen0(poll, recv_handler, context) { log::error!("{:?}", e); } + drop(_waker); worker.stop_all(); })?; Ok(()) } -pub fn main_udp_listen0(index: usize, mut recv_handler: H, context: Context) -> io::Result<()> +pub fn main_udp_listen0(mut poll: Poll, mut recv_handler: H, context: Context) -> io::Result<()> where H: RecvChannelHandler, { let mut buf = [0; BUFFER_SIZE]; - let udp_socket = &context.main_udp_socket[index]; + let mut udps = Vec::with_capacity(context.main_udp_socket.len()); + + for (index, udp) in context.main_udp_socket.iter().enumerate() { + let udp_socket = udp.try_clone()?; + udp_socket.set_nonblocking(true)?; + let mut mio_udp = UdpSocket::from_std(udp_socket); + poll.registry() + .register(&mut mio_udp, Token(index + 1), Interest::READABLE)?; + udps.push(mio_udp); + } + + let mut events = Events::with_capacity(udps.len()); loop { - match udp_socket.recv_from(&mut buf) { - Ok((len, addr)) => { - if &buf[..len] == b"stop" { - if context.is_stop() { - return Ok(()); + poll.poll(&mut events, None)?; + for x in events.iter() { + let index = match x.token() { + NOTIFY => return Ok(()), + Token(index) => index - 1, + }; + loop { + match udps[index].recv_from(&mut buf) { + Ok((len, addr)) => { + recv_handler.handle( + &mut buf[..len], + RouteKey::new(false, index, addr), + &context, + ); + } + Err(e) => { + if e.kind() == io::ErrorKind::WouldBlock { + break; + } + log::error!("main_udp_listen_{}={:?}", index, e); } } - recv_handler.handle(&mut buf[..len], RouteKey::new(false, index, addr), &context); - } - Err(e) => { - log::error!("main_udp_listen0={:?}", e); } } } } +// /// 用recvmmsg没什么帮助,这里记录下,以下是完整代码 +// #[cfg(unix)] +// pub fn main_udp_listen0(index: usize, mut recv_handler: H, context: Context) -> io::Result<()> +// where +// H: RecvChannelHandler, +// { +// use libc::{c_uint, mmsghdr, sockaddr_storage, socklen_t, timespec}; +// use std::os::fd::AsRawFd; +// +// let udp_socket = context.main_udp_socket[index].try_clone()?; +// let fd = udp_socket.as_raw_fd(); +// const MAX_MESSAGES: usize = 16; +// let mut iov: [libc::iovec; MAX_MESSAGES] = unsafe { std::mem::zeroed() }; +// let mut buf: [[u8; BUFFER_SIZE]; MAX_MESSAGES] = [[0; BUFFER_SIZE]; MAX_MESSAGES]; +// let mut msgs: [mmsghdr; MAX_MESSAGES] = unsafe { std::mem::zeroed() }; +// let mut addrs: [sockaddr_storage; MAX_MESSAGES] = unsafe { std::mem::zeroed() }; +// for i in 0..MAX_MESSAGES { +// iov[i].iov_base = buf[i].as_mut_ptr() as *mut libc::c_void; +// iov[i].iov_len = BUFFER_SIZE; +// msgs[i].msg_hdr.msg_iov = &mut iov[i]; +// msgs[i].msg_hdr.msg_iovlen = 1; +// msgs[i].msg_hdr.msg_name = &mut addrs[i] as *const _ as *mut libc::c_void; +// msgs[i].msg_hdr.msg_namelen = std::mem::size_of::() as socklen_t; +// } +// let mut time: timespec = unsafe { std::mem::zeroed() }; +// loop { +// if context.is_stop() { +// return Ok(()); +// } +// let res = +// unsafe { libc::recvmmsg(fd, msgs.as_mut_ptr(), MAX_MESSAGES as c_uint, 0, &mut time) }; +// if res == -1 { +// log::error!("main_udp_listen_{}={:?}", index, io::Error::last_os_error()); +// continue; +// } +// +// let nmsgs = res as usize; +// for i in 0..nmsgs { +// let msg = &mut buf[i][0..msgs[i].msg_len as usize]; +// let addr = sockaddr_to_socket_addr(&addrs[i], msgs[i].msg_hdr.msg_namelen); +// if msg == b"stop" { +// if context.is_stop() { +// return Ok(()); +// } +// } +// recv_handler.handle(msg, RouteKey::new(false, index, addr), &context); +// } +// } +// } +// +// #[cfg(unix)] +// fn sockaddr_to_socket_addr(addr: &libc::sockaddr_storage, _len: libc::socklen_t) -> SocketAddr { +// match addr.ss_family as libc::c_int { +// libc::AF_INET => { +// let addr_in = unsafe { *(addr as *const _ as *const libc::sockaddr_in) }; +// let ip = u32::from_be(addr_in.sin_addr.s_addr); +// let port = u16::from_be(addr_in.sin_port); +// SocketAddr::V4(std::net::SocketAddrV4::new(Ipv4Addr::from(ip), port)) +// } +// libc::AF_INET6 => { +// let addr_in6 = unsafe { *(addr as *const _ as *const libc::sockaddr_in6) }; +// let ip = std::net::Ipv6Addr::from(addr_in6.sin6_addr.s6_addr); +// let port = u16::from_be(addr_in6.sin6_port); +// SocketAddr::V6(std::net::SocketAddrV6::new(ip, port, 0, 0)) +// } +// _ => panic!("Unsupported address family"), +// } +// }