[mio] 主通道改为异步

This commit is contained in:
lubeilin
2024-03-10 14:11:06 +08:00
parent 03529bc339
commit b2b58a23c1
3 changed files with 194 additions and 90 deletions
+67 -44
View File
@@ -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<Route> {
fn get_route_by_id(&self, index: usize, id: &Ipv4Addr) -> io::Result<Route> {
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<Route> {
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<Ipv4Addr> {
let table = self.route_table.read();
for (k, (_, v)) in table.iter() {
+15 -5
View File
@@ -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());
+112 -41
View File
@@ -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<H>(
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<H>(
index: usize,
stop_manager: StopManager,
recv_handler: H,
context: Context,
@@ -154,54 +144,135 @@ fn main_udp_listen<H>(
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<H>(index: usize, mut recv_handler: H, context: Context) -> io::Result<()>
pub fn main_udp_listen0<H>(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<H>(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::<sockaddr_storage>() 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"),
// }
// }