use crate::RuntimeError;
use crate::runtime_state::LifecycleSignals;
use std::future::Future;
use std::net::SocketAddr;
use std::sync::Arc;
const MAX_DATAGRAM: usize = 65535;
#[derive(Debug)]
pub struct UdpSocket {
inner: tokio::net::UdpSocket,
}
impl UdpSocket {
pub async fn bind(addr: &str) -> Result<Self, RuntimeError> {
let inner = tokio::net::UdpSocket::bind(addr).await?;
Ok(Self { inner })
}
pub async fn connect(&self, addr: &str) -> Result<(), RuntimeError> {
self.inner.connect(addr).await?;
Ok(())
}
pub async fn send_to<A>(&self, datagram: &[u8], target: A) -> Result<usize, RuntimeError>
where
A: tokio::net::ToSocketAddrs,
{
let bytes_sent = self.inner.send_to(datagram, target).await?;
Ok(bytes_sent)
}
pub async fn recv_from(
&self,
recv_buf: &mut [u8],
) -> Result<(usize, SocketAddr), RuntimeError> {
let (bytes_read, addr) = self.inner.recv_from(recv_buf).await?;
Ok((bytes_read, addr))
}
pub async fn send(&self, datagram: &[u8]) -> Result<usize, RuntimeError> {
let bytes_sent = self.inner.send(datagram).await?;
Ok(bytes_sent)
}
pub async fn recv(&self, recv_buf: &mut [u8]) -> Result<usize, RuntimeError> {
let bytes_read = self.inner.recv(recv_buf).await?;
Ok(bytes_read)
}
pub fn local_addr(&self) -> Result<SocketAddr, RuntimeError> {
let addr = self.inner.local_addr()?;
Ok(addr)
}
}
pub async fn serve_udp<F, Fut>(addr: &str, handler: F) -> Result<(), RuntimeError>
where
F: Fn(Vec<u8>, SocketAddr, Arc<UdpSocket>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), RuntimeError>> + Send,
{
let socket = UdpSocket::bind(addr).await?;
serve_udp_on(socket, handler).await
}
pub async fn serve_udp_on<F, Fut>(socket: UdpSocket, handler: F) -> Result<(), RuntimeError>
where
F: Fn(Vec<u8>, SocketAddr, Arc<UdpSocket>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), RuntimeError>> + Send,
{
let signals = LifecycleSignals::current();
let socket = Arc::new(socket);
recv_loop(&socket, &signals, &handler).await
}
async fn recv_loop<F, Fut>(
socket: &Arc<UdpSocket>,
signals: &LifecycleSignals,
handler: &F,
) -> Result<(), RuntimeError>
where
F: Fn(Vec<u8>, SocketAddr, Arc<UdpSocket>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), RuntimeError>> + Send,
{
let mut buf = vec![0u8; MAX_DATAGRAM].into_boxed_slice();
let stop = signals.wait();
tokio::pin!(stop);
loop {
let received = tokio::select! {
biased;
() = &mut stop => return Ok(()),
result = socket.recv_from(&mut buf) => result,
};
match classify_receive(received)? {
None => {}
Some((len, addr)) => dispatch(&buf[..len], addr, socket, handler).await,
}
}
}
fn classify_receive(
result: Result<(usize, SocketAddr), RuntimeError>,
) -> Result<Option<(usize, SocketAddr)>, RuntimeError> {
match result {
Ok(received) => Ok(Some(received)),
Err(error) if crate::error::is_transient_datagram_error(&error) => {
tracing::debug!(%error, "udp recv: transient error, continuing");
Ok(None)
}
Err(error) => Err(error),
}
}
async fn dispatch<F, Fut>(datagram: &[u8], addr: SocketAddr, socket: &Arc<UdpSocket>, handler: &F)
where
F: Fn(Vec<u8>, SocketAddr, Arc<UdpSocket>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), RuntimeError>> + Send,
{
let result = handler(datagram.to_vec(), addr, Arc::clone(socket)).await;
super::accept::report_handler_error("udp", result);
}