use std::{collections::HashMap, num::NonZeroUsize, path::PathBuf, sync::Arc};
use tokio::sync::{Mutex, mpsc};
use tracing::{debug, error, info, warn};
use crate::{buffer::Buffer, endpoint::EndpointStream, shutdown::Shutdown};
const SESSION_DEPTH: usize = 64;
const PENDING_SESSIONS: usize = 16;
#[derive(Clone)]
pub enum Socket {
Udp(Arc<tokio::net::UdpSocket>),
Unix(Arc<tokio::net::UnixDatagram>),
}
#[derive(Clone, PartialEq, Eq, Hash)]
pub enum Peer {
Inet(std::net::SocketAddr),
Path(PathBuf),
}
impl Socket {
async fn recv_from(&self, buf: &mut [u8]) -> std::io::Result<(usize, Option<Peer>)> {
match self {
Socket::Udp(socket) => {
let (n, peer) = socket.recv_from(buf).await?;
Ok((n, Some(Peer::Inet(peer))))
}
Socket::Unix(socket) => {
let (n, peer) = socket.recv_from(buf).await?;
Ok((n, peer.as_pathname().map(|p| Peer::Path(p.to_path_buf()))))
}
}
}
async fn send_to(&self, buf: &[u8], peer: &Peer) -> std::io::Result<usize> {
match (self, peer) {
(Socket::Udp(socket), Peer::Inet(peer)) => socket.send_to(buf, *peer).await,
(Socket::Unix(socket), Peer::Path(peer)) => socket.send_to(buf, peer).await,
_ => Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"the peer address does not belong to this socket",
)),
}
}
}
impl std::fmt::Display for Peer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Peer::Inet(peer) => write!(f, "{peer}"),
Peer::Path(peer) => write!(f, "{}", peer.display()),
}
}
}
pub(super) fn demux(socket: Socket, max: NonZeroUsize, buffer: usize, shutdown: Shutdown) -> Demux {
let (peers, incoming) = mpsc::channel(PENDING_SESSIONS);
tokio::spawn(demultiplex(socket, peers, max, buffer, shutdown));
Demux { incoming }
}
pub struct Demux {
incoming: mpsc::Receiver<(EndpointStream, String)>,
}
impl Demux {
pub async fn accept(&mut self) -> std::io::Result<(EndpointStream, String)> {
self.incoming.recv().await.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"the datagram receive loop stopped",
)
})
}
}
pub struct Session {
peer: Peer,
socket: Socket,
queue: Mutex<mpsc::Receiver<Vec<u8>>>,
}
impl Session {
fn new(peer: Peer, socket: Socket, queue: mpsc::Receiver<Vec<u8>>) -> Self {
Self {
peer,
socket,
queue: Mutex::new(queue),
}
}
pub(super) async fn recv(&self, buf: &mut [u8]) -> std::io::Result<Option<usize>> {
let Some(message) = self.queue.lock().await.recv().await else {
return Ok(None);
};
let n = message.len().min(buf.len());
buf[..n].copy_from_slice(&message[..n]);
Ok(Some(n))
}
pub(super) async fn send(&self, buf: &[u8]) -> std::io::Result<usize> {
self.socket.send_to(buf, &self.peer).await
}
}
async fn demultiplex(
socket: Socket,
peers: mpsc::Sender<(EndpointStream, String)>,
max: NonZeroUsize,
buffer: usize,
mut shutdown: Shutdown,
) {
let mut sessions: HashMap<Peer, mpsc::Sender<Vec<u8>>> = HashMap::new();
let mut buf = Buffer::new(buffer);
let mut dropped = 0u64;
loop {
let received = tokio::select! {
biased;
received = socket.recv_from(&mut buf) => received,
() = shutdown.recv() => break,
};
let (n, peer) = match received {
Ok(received) => received,
Err(e) => {
error!(error = %e, "receiving datagram failed, stopping");
break;
}
};
let Some(peer) = peer else {
drop_datagram(&mut dropped, None, "the sender has no address to reply to");
continue;
};
let delivered = sessions
.get(&peer)
.map(|session| session.try_send(buf[..n].to_vec()));
match delivered {
Some(Ok(())) => continue,
Some(Err(mpsc::error::TrySendError::Full(_))) => {
drop_datagram(&mut dropped, Some(&peer), "session is not keeping up");
continue;
}
Some(Err(mpsc::error::TrySendError::Closed(_))) => {
sessions.remove(&peer);
}
None => {}
}
if sessions.len() >= max.get() {
sessions.retain(|_, session| !session.is_closed());
if sessions.len() >= max.get() {
drop_datagram(
&mut dropped,
Some(&peer),
"ceiling reached; add a timeout stage to reclaim quiet sessions",
);
continue;
}
}
let (inbox, queue) = mpsc::channel(SESSION_DEPTH);
let session = Session::new(peer.clone(), socket.clone(), queue);
match peers.try_send((EndpointStream::datagram_session(session), peer.to_string())) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
drop_datagram(
&mut dropped,
Some(&peer),
"relay is not accepting quickly enough",
);
continue;
}
Err(mpsc::error::TrySendError::Closed(_)) => break,
}
let _ = inbox.try_send(buf[..n].to_vec());
sessions.insert(peer.clone(), inbox);
info!(%peer, sessions = sessions.len(), "new sender");
}
debug!(dropped, "datagram receive loop finished");
}
fn drop_datagram(dropped: &mut u64, peer: Option<&Peer>, reason: &'static str) {
*dropped += 1;
let peer = peer.map_or_else(|| "unnamed".to_owned(), Peer::to_string);
if *dropped == 1 {
warn!(
peer,
reason, "dropping datagram; further drops are logged at debug"
);
} else {
debug!(peer, reason, "dropping datagram");
}
}