use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::UdpSocket;
use super::frame::sizes;
pub const DEFAULT_RECV_BUFFER_SIZE: usize = 65535;
#[derive(Debug)]
pub struct NomadSocket {
socket: Arc<UdpSocket>,
recv_buffer: Vec<u8>,
max_payload_size: usize,
}
impl NomadSocket {
pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
let socket = UdpSocket::bind(addr).await?;
Ok(Self::from_socket(socket))
}
pub fn from_socket(socket: UdpSocket) -> Self {
Self {
socket: Arc::new(socket),
recv_buffer: vec![0u8; DEFAULT_RECV_BUFFER_SIZE],
max_payload_size: sizes::DEFAULT_MAX_PAYLOAD,
}
}
pub fn set_max_payload_size(&mut self, size: usize) {
self.max_payload_size = size;
}
pub fn max_payload_size(&self) -> usize {
self.max_payload_size
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
pub async fn connect(&self, addr: SocketAddr) -> io::Result<()> {
self.socket.connect(addr).await
}
pub async fn send_to(&self, data: &[u8], addr: SocketAddr) -> io::Result<usize> {
self.socket.send_to(data, addr).await
}
pub async fn send(&self, data: &[u8]) -> io::Result<usize> {
self.socket.send(data).await
}
pub async fn recv_from(&mut self) -> io::Result<(&[u8], SocketAddr)> {
let (len, addr) = self.socket.recv_from(&mut self.recv_buffer).await?;
Ok((&self.recv_buffer[..len], addr))
}
pub async fn recv(&mut self) -> io::Result<&[u8]> {
let len = self.socket.recv(&mut self.recv_buffer).await?;
Ok(&self.recv_buffer[..len])
}
pub fn try_recv_from(&mut self) -> io::Result<Option<(usize, SocketAddr)>> {
match self.socket.try_recv_from(&mut self.recv_buffer) {
Ok((len, addr)) => Ok(Some((len, addr))),
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => Ok(None),
Err(e) => Err(e),
}
}
pub fn recv_data(&self, len: usize) -> &[u8] {
&self.recv_buffer[..len]
}
pub fn inner(&self) -> &UdpSocket {
&self.socket
}
pub fn socket_arc(&self) -> Arc<UdpSocket> {
Arc::clone(&self.socket)
}
pub fn max_frame_size(&self) -> usize {
self.max_payload_size + sizes::DATA_FRAME_HEADER_SIZE + sizes::AEAD_TAG_SIZE
}
}
#[derive(Debug, Clone)]
pub struct NomadSocketBuilder {
recv_buffer_size: usize,
max_payload_size: usize,
}
impl Default for NomadSocketBuilder {
fn default() -> Self {
Self::new()
}
}
impl NomadSocketBuilder {
pub fn new() -> Self {
Self {
recv_buffer_size: DEFAULT_RECV_BUFFER_SIZE,
max_payload_size: sizes::DEFAULT_MAX_PAYLOAD,
}
}
pub fn recv_buffer_size(mut self, size: usize) -> Self {
self.recv_buffer_size = size;
self
}
pub fn max_payload_size(mut self, size: usize) -> Self {
self.max_payload_size = size;
self
}
pub async fn bind(self, addr: SocketAddr) -> io::Result<NomadSocket> {
let socket = UdpSocket::bind(addr).await?;
Ok(self.from_socket(socket))
}
pub fn from_socket(self, socket: UdpSocket) -> NomadSocket {
NomadSocket {
socket: Arc::new(socket),
recv_buffer: vec![0u8; self.recv_buffer_size],
max_payload_size: self.max_payload_size,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_socket_bind() {
let socket = NomadSocket::bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let addr = socket.local_addr().unwrap();
assert!(addr.port() != 0);
}
#[tokio::test]
async fn test_socket_send_recv() {
let mut server = NomadSocket::bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let server_addr = server.local_addr().unwrap();
let client = NomadSocket::bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let data = b"hello NOMAD";
client.send_to(data, server_addr).await.unwrap();
let (received, from) = server.recv_from().await.unwrap();
assert_eq!(received, data);
assert_eq!(from, client.local_addr().unwrap());
}
#[tokio::test]
async fn test_socket_connected() {
let mut server = NomadSocket::bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let server_addr = server.local_addr().unwrap();
let client = NomadSocket::bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
client.connect(server_addr).await.unwrap();
let data = b"connected send";
client.send(data).await.unwrap();
let (received, _) = server.recv_from().await.unwrap();
assert_eq!(received, data);
}
#[test]
fn test_socket_builder() {
let builder = NomadSocketBuilder::new()
.recv_buffer_size(4096)
.max_payload_size(1400);
assert_eq!(builder.recv_buffer_size, 4096);
assert_eq!(builder.max_payload_size, 1400);
}
#[tokio::test]
async fn test_max_frame_size() {
let socket = NomadSocketBuilder::new()
.max_payload_size(1200)
.bind("127.0.0.1:0".parse().unwrap())
.await
.unwrap();
let expected = 1200 + sizes::DATA_FRAME_HEADER_SIZE + sizes::AEAD_TAG_SIZE;
assert_eq!(socket.max_frame_size(), expected);
}
}