use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use async_trait::async_trait;
use socket2::SockRef;
use tokio::net::UdpSocket;
use tracing::{debug, warn};
#[async_trait]
pub trait Transport: Send + Sync + 'static {
async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)>;
async fn send_to(&self, buf: &[u8], dst: &SocketAddr) -> io::Result<usize>;
fn local_addr(&self) -> io::Result<SocketAddr>;
}
#[derive(Clone, Debug)]
pub struct UdpTransport(Arc<UdpSocket>);
impl UdpTransport {
pub fn new(socket: Arc<UdpSocket>) -> Self {
UdpTransport(socket)
}
pub async fn bind(
addr: SocketAddr,
recv_buffer_size: Option<usize>,
send_buffer_size: Option<usize>,
) -> io::Result<Self> {
let socket = UdpSocket::bind(addr).await?;
set_socket_buffers(&socket, recv_buffer_size, send_buffer_size);
Ok(UdpTransport(Arc::new(socket)))
}
pub fn socket(&self) -> &UdpSocket {
&self.0
}
}
fn set_socket_buffers(
socket: &UdpSocket,
recv_buffer_size: Option<usize>,
send_buffer_size: Option<usize>,
) {
let sock = SockRef::from(socket);
if let Some(size) = recv_buffer_size {
match sock.set_recv_buffer_size(size) {
Ok(()) => match sock.recv_buffer_size() {
Ok(actual) => debug!(
"gossip socket SO_RCVBUF: requested {size} B, OS granted {actual} B \
(raise net.core.rmem_max if a larger buffer is needed)"
),
Err(e) => debug!("could not read back SO_RCVBUF: {e}"),
},
Err(e) => warn!("failed to set gossip socket SO_RCVBUF to {size} B: {e}"),
}
}
if let Some(size) = send_buffer_size {
match sock.set_send_buffer_size(size) {
Ok(()) => match sock.send_buffer_size() {
Ok(actual) => debug!(
"gossip socket SO_SNDBUF: requested {size} B, OS granted {actual} B \
(raise net.core.wmem_max if a larger buffer is needed)"
),
Err(e) => debug!("could not read back SO_SNDBUF: {e}"),
},
Err(e) => warn!("failed to set gossip socket SO_SNDBUF to {size} B: {e}"),
}
}
}
#[async_trait]
impl Transport for UdpTransport {
async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
self.0.recv_from(buf).await
}
async fn send_to(&self, buf: &[u8], dst: &SocketAddr) -> io::Result<usize> {
self.0.send_to(buf, dst).await
}
fn local_addr(&self) -> io::Result<SocketAddr> {
self.0.local_addr()
}
}
pub use in_memory::{InMemoryNetwork, InMemoryTransport};
mod in_memory {
use super::*;
use std::collections::HashMap;
use parking_lot::Mutex;
use tokio::sync::mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender};
use tokio::sync::Mutex as AsyncMutex;
type Datagram = (SocketAddr, Vec<u8>);
#[derive(Clone, Default, Debug)]
pub struct InMemoryNetwork {
routes: Arc<Mutex<HashMap<SocketAddr, UnboundedSender<Datagram>>>>,
}
impl InMemoryNetwork {
pub fn new() -> Self {
InMemoryNetwork::default()
}
pub fn bind(&self, addr: SocketAddr) -> InMemoryTransport {
let (tx, rx) = unbounded_channel();
self.routes.lock().insert(addr, tx);
InMemoryTransport {
network: self.clone(),
addr,
rx: AsyncMutex::new(rx),
}
}
}
#[derive(Debug)]
pub struct InMemoryTransport {
network: InMemoryNetwork,
addr: SocketAddr,
rx: AsyncMutex<UnboundedReceiver<Datagram>>,
}
#[async_trait]
impl Transport for InMemoryTransport {
async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
let (src, bytes) = self.rx.lock().await.recv().await.ok_or_else(|| {
io::Error::new(io::ErrorKind::BrokenPipe, "in-memory network closed")
})?;
let n = bytes.len().min(buf.len());
buf[..n].copy_from_slice(&bytes[..n]);
Ok((n, src))
}
async fn send_to(&self, buf: &[u8], dst: &SocketAddr) -> io::Result<usize> {
if let Some(tx) = self.network.routes.lock().get(dst) {
let _ = tx.send((self.addr, buf.to_vec()));
}
Ok(buf.len())
}
fn local_addr(&self) -> io::Result<SocketAddr> {
Ok(self.addr)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn in_memory_delivers_between_endpoints() {
let net = InMemoryNetwork::new();
let a: SocketAddr = "127.0.0.1:1".parse().unwrap();
let b: SocketAddr = "127.0.0.1:2".parse().unwrap();
let ta = net.bind(a);
let tb = net.bind(b);
ta.send_to(b"hello", &b).await.unwrap();
let mut buf = [0u8; 16];
let (n, src) = tb.recv_from(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"hello");
assert_eq!(src, a);
assert_eq!(tb.local_addr().unwrap(), b);
}
#[tokio::test]
async fn in_memory_drops_datagram_to_unknown_address() {
let net = InMemoryNetwork::new();
let a: SocketAddr = "127.0.0.1:1".parse().unwrap();
let ta = net.bind(a);
let unknown: SocketAddr = "127.0.0.1:9".parse().unwrap();
let n = ta.send_to(b"lost", &unknown).await.unwrap();
assert_eq!(n, 4);
}
}