use dashmap::DashMap;
use std::convert::TryFrom;
use std::error::Error;
use std::future::Future;
use std::io;
use std::marker::PhantomData;
use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use std::time::Duration;
use crate::peer::{UDPPeer, UdpPeer, UdpReader};
use bytes::BytesMut;
use socket2::{Domain, Protocol, Socket, Type};
use tokio::net::UdpSocket;
use tokio::sync::mpsc::unbounded_channel;
pub const BUFF_MAX_SIZE: usize = 4096;
pub const DEFAULT_BUF_SIZE: usize = 1784 * 10000;
pub struct UdpContext {
#[allow(dead_code)]
pub id: usize,
recv: Arc<UdpSocket>,
pub peers: DashMap<SocketAddr, UDPPeer>,
}
pub struct UdpServer<I, T> {
addr: SocketAddr,
input: Arc<I>,
_ph: PhantomData<T>,
clean_sec: Option<u64>,
buf_size: usize,
}
impl<I, R, T> UdpServer<I, T>
where
I: Fn(UDPPeer, UdpReader, T) -> R + Send + Sync + 'static,
R: Future<Output = Result<(), Box<dyn Error>>> + Send + 'static,
T: Sync + Send + Clone + 'static,
{
pub fn new<A: ToSocketAddrs>(addr: A, input: I) -> io::Result<Self> {
let addr = resolve_single_addr(&addr)?;
Ok(UdpServer {
addr,
input: Arc::new(input),
_ph: Default::default(),
clean_sec: None,
buf_size: DEFAULT_BUF_SIZE,
})
}
#[inline]
pub fn set_buffer_size(mut self, size: usize) -> UdpServer<I, T> {
assert!(size > 0, "buffer size must be greater than 0");
self.buf_size = size;
self
}
#[inline]
pub fn set_peer_timeout_sec(mut self, sec: u64) -> UdpServer<I, T> {
assert!(sec > 0, "timeout must be greater than 0");
self.clean_sec = Some(sec);
self
}
#[inline]
pub async fn start(&self, inner: T) -> io::Result<()> {
let udp_list = create_udp_socket_list(&self.addr, get_cpu_count(), self.buf_size)?;
let udp_contexts: Vec<Arc<UdpContext>> = udp_list
.into_iter()
.enumerate()
.map(|(id, socket)| {
Arc::new(UdpContext {
id,
recv: Arc::new(socket),
peers: Default::default(),
})
})
.collect();
let need_check_timeout = {
if let Some(clean_sec) = self.clean_sec {
let clean_sec = clean_sec as i64;
let contexts = udp_contexts.clone();
tokio::spawn(async move {
loop {
let current = chrono::Utc::now().timestamp();
for context in contexts.iter() {
context.peers.iter().for_each(|entry| {
let peer = entry.value();
if current - peer.get_last_recv_sec() > clean_sec {
peer.close();
}
});
}
tokio::time::sleep(Duration::from_secs(1)).await
}
});
true
} else {
false
}
};
let (tx, mut rx) = unbounded_channel();
for (index, udp_listen) in udp_contexts.iter().enumerate() {
let create_peer_tx = tx.clone();
let udp_context = udp_listen.clone();
tokio::spawn(async move {
log::debug!("start udp listen:{index}");
let mut buf = BytesMut::zeroed(BUFF_MAX_SIZE);
loop {
match udp_context.recv.recv_from(&mut buf).await {
Ok((size, addr)) => {
let data = buf.split_to(size).freeze();
buf = BytesMut::zeroed(BUFF_MAX_SIZE);
let peer = {
udp_context
.peers
.entry(addr)
.or_insert_with(|| {
let (peer, reader) =
UdpPeer::new(index, udp_context.recv.clone(), addr);
log::trace!("create udp listen:{index} udp peer:{addr}");
if let Err(err) =
create_peer_tx.send((peer.clone(), reader, index, addr))
{
panic!("create_peer_tx err:{}", err);
}
peer
})
.clone()
};
if need_check_timeout {
if let Err(err) = peer.push_data_and_update_instant(data).await {
log::error!("peer push data and update instant is error:{err}");
}
} else if let Err(err) = peer.push_data(data) {
log::error!("peer push data is error:{err}");
}
}
Err(err) => {
log::trace!("udp:{index} recv_from error:{err}");
}
}
}
});
}
drop(tx);
while let Some((peer, reader, index, addr)) = rx.recv().await {
let inner = inner.clone();
let input_fn = self.input.clone();
let context = udp_contexts.get(index).expect("not found context").clone();
tokio::spawn(async move {
if let Err(err) = (input_fn)(peer, reader, inner).await {
log::error!("udp input error:{err}")
}
context.peers.remove(&addr);
});
}
Ok(())
}
}
fn resolve_single_addr<A: ToSocketAddrs>(addr: &A) -> io::Result<SocketAddr> {
let mut addrs = addr.to_socket_addrs()?;
let addr = match addrs.next() {
Some(addr) => addr,
None => return Err(io::Error::other("no socket addresses could be resolved")),
};
if addrs.next().is_some() {
return Err(io::Error::other("more than one address resolved"));
}
Ok(addr)
}
fn create_udp_socket(addr: &SocketAddr, buf_size: usize) -> io::Result<std::net::UdpSocket> {
let domain = if addr.is_ipv4() {
Domain::IPV4
} else if addr.is_ipv6() {
Domain::IPV6
} else {
return Err(io::Error::other("not address AF_INET"));
};
let socket = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
socket.set_reuse_address(true)?;
#[cfg(not(target_os = "windows"))]
socket.set_reuse_port(true)?;
socket.bind(&(*addr).into())?;
socket.set_send_buffer_size(buf_size)?;
socket.set_recv_buffer_size(buf_size)?;
Ok(socket.into())
}
fn create_async_udp_socket(addr: &SocketAddr, buf_size: usize) -> io::Result<UdpSocket> {
let std_sock = create_udp_socket(addr, buf_size)?;
std_sock.set_nonblocking(true)?;
let sock = UdpSocket::try_from(std_sock)?;
Ok(sock)
}
fn create_udp_socket_list(
addr: &SocketAddr,
listen_count: usize,
buf_size: usize,
) -> io::Result<Vec<UdpSocket>> {
log::debug!("cpus:{listen_count}");
let mut listens = Vec::with_capacity(listen_count);
for _ in 0..listen_count {
let sock = create_async_udp_socket(addr, buf_size)?;
listens.push(sock);
}
Ok(listens)
}
#[cfg(not(target_os = "windows"))]
fn get_cpu_count() -> usize {
num_cpus::get()
}
#[cfg(target_os = "windows")]
fn get_cpu_count() -> usize {
1
}