use core::future::Future;
use core::net::{Ipv4Addr, SocketAddrV4};
use core::pin::Pin;
use core::task::{Context, Poll};
use core::time::Duration;
use std::net::{IpAddr, SocketAddr};
use tokio::io::ReadBuf;
use tokio::net::UdpSocket;
use crate::transport::{
ChannelFactory, IoErrorKind, MpscRecv, MpscSend, OneshotCancelled, OneshotRecv, OneshotSend,
ReceivedDatagram, SocketOptions, Timer, TransportError, TransportFactory, TransportSocket,
UnboundedRecv, UnboundedSend,
};
#[derive(Debug, Default, Clone, Copy)]
pub struct TokioTransport;
#[derive(Debug)]
pub struct TokioSocket {
inner: UdpSocket,
}
impl TokioSocket {
#[allow(dead_code)] pub(crate) fn multicast_loop_v4(&self) -> Result<bool, TransportError> {
self.inner.multicast_loop_v4().map_err(|e| map_io_error(&e))
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct TokioTimer;
#[derive(Debug, Default, Clone, Copy)]
pub struct TokioSpawner;
pub struct TokioBindFuture {
addr: SocketAddrV4,
options: SocketOptions,
}
impl Future for TokioBindFuture {
type Output = Result<TokioSocket, TransportError>;
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
let addr = self.addr;
let options = self.options;
Poll::Ready(bind_with_options(addr, options).map_err(|e| map_io_error(&e)))
}
}
impl TransportFactory for TokioTransport {
type Socket = TokioSocket;
type BindFuture<'a> = TokioBindFuture;
fn bind<'a>(&'a self, addr: SocketAddrV4, options: &'a SocketOptions) -> Self::BindFuture<'a> {
TokioBindFuture {
addr,
options: *options,
}
}
}
pub struct SendTo<'a> {
socket: &'a UdpSocket,
buf: &'a [u8],
target: SocketAddr,
}
impl Future for SendTo<'_> {
type Output = Result<(), TransportError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match self.socket.poll_send_to(cx, self.buf, self.target) {
Poll::Pending => Poll::Pending,
Poll::Ready(Ok(_n)) => Poll::Ready(Ok(())),
Poll::Ready(Err(e)) => Poll::Ready(Err(map_io_error(&e))),
}
}
}
pub struct RecvFrom<'a> {
socket: &'a UdpSocket,
buf: &'a mut [u8],
}
impl Future for RecvFrom<'_> {
type Output = Result<ReceivedDatagram, TransportError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let me = self.get_mut();
let mut read_buf = ReadBuf::new(me.buf);
match me.socket.poll_recv_from(cx, &mut read_buf) {
Poll::Pending => Poll::Pending,
Poll::Ready(Err(e)) => Poll::Ready(Err(map_io_error(&e))),
Poll::Ready(Ok(src)) => {
let n = read_buf.filled().len();
let source = match src {
SocketAddr::V4(v4) => v4,
SocketAddr::V6(_) => return Poll::Ready(Err(TransportError::Unsupported)),
};
Poll::Ready(Ok(ReceivedDatagram {
bytes_received: n,
source,
truncated: false,
}))
}
}
}
}
impl TransportSocket for TokioSocket {
type SendFuture<'a> = SendTo<'a>;
type RecvFuture<'a> = RecvFrom<'a>;
fn send_to<'a>(&'a self, buf: &'a [u8], target: SocketAddrV4) -> Self::SendFuture<'a> {
SendTo {
socket: &self.inner,
buf,
target: SocketAddr::V4(target),
}
}
fn recv_from<'a>(&'a self, buf: &'a mut [u8]) -> Self::RecvFuture<'a> {
RecvFrom {
socket: &self.inner,
buf,
}
}
fn local_addr(&self) -> Result<SocketAddrV4, TransportError> {
match self.inner.local_addr().map_err(|e| map_io_error(&e))? {
SocketAddr::V4(v4) => Ok(v4),
SocketAddr::V6(_) => Err(TransportError::Unsupported),
}
}
fn join_multicast_v4(&self, group: Ipv4Addr, iface: Ipv4Addr) -> Result<(), TransportError> {
self.inner
.join_multicast_v4(group, iface)
.map_err(|e| map_io_error(&e))
}
fn leave_multicast_v4(&self, group: Ipv4Addr, iface: Ipv4Addr) -> Result<(), TransportError> {
self.inner
.leave_multicast_v4(group, iface)
.map_err(|e| map_io_error(&e))
}
}
pub struct TokioSleep {
inner: tokio::time::Sleep,
}
impl Future for TokioSleep {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let inner = unsafe { self.map_unchecked_mut(|s| &mut s.inner) };
inner.poll(cx).map(|()| ())
}
}
impl Timer for TokioTimer {
type SleepFuture<'a> = TokioSleep;
fn sleep(&self, duration: Duration) -> Self::SleepFuture<'_> {
TokioSleep {
inner: tokio::time::sleep(duration),
}
}
}
struct PanicLoggingFut<F> {
inner: F,
}
impl<F: Future<Output = ()>> Future for PanicLoggingFut<F> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let inner = unsafe { self.map_unchecked_mut(|s| &mut s.inner) };
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| inner.poll(cx))) {
Ok(poll) => poll,
Err(payload) => {
let msg = panic_payload_str(&payload);
crate::log::error!(
panic_message = msg,
"spawned task panicked; channels will close",
);
Poll::Ready(())
}
}
}
}
impl crate::transport::Spawner for TokioSpawner {
fn spawn(&self, future: impl Future<Output = ()> + Send + 'static) {
drop(tokio::spawn(PanicLoggingFut { inner: future }));
}
}
fn panic_payload_str(payload: &std::boxed::Box<dyn std::any::Any + Send>) -> &str {
if let Some(s) = payload.downcast_ref::<&'static str>() {
s
} else if let Some(s) = payload.downcast_ref::<std::string::String>() {
s.as_str()
} else {
"<non-string panic payload>"
}
}
fn bind_with_options(addr: SocketAddrV4, options: SocketOptions) -> std::io::Result<TokioSocket> {
let raw = socket2::Socket::new(
socket2::Domain::IPV4,
socket2::Type::DGRAM,
Some(socket2::Protocol::UDP),
)?;
if options.reuse_address {
raw.set_reuse_address(true)?;
}
#[cfg(unix)]
if options.reuse_port {
raw.set_reuse_port(true)?;
}
if let Some(iface) = options.multicast_if_v4 {
raw.set_multicast_if_v4(&iface)?;
}
if let Some(loop_v4) = options.multicast_loop_v4 {
raw.set_multicast_loop_v4(loop_v4)?;
}
let bind_addr = SocketAddr::new(IpAddr::V4(*addr.ip()), addr.port());
raw.bind(&bind_addr.into())?;
raw.set_nonblocking(true)?;
let std_sock: std::net::UdpSocket = raw.into();
let inner = UdpSocket::from_std(std_sock)?;
Ok(TokioSocket { inner })
}
fn map_io_error(e: &std::io::Error) -> TransportError {
use std::io::ErrorKind as K;
let kind = e.kind();
let mapped = match kind {
K::AddrInUse => TransportError::AddressInUse,
K::Unsupported => TransportError::Unsupported,
K::TimedOut => TransportError::Io(IoErrorKind::TimedOut),
K::Interrupted => TransportError::Io(IoErrorKind::Interrupted),
K::PermissionDenied => TransportError::Io(IoErrorKind::PermissionDenied),
K::ConnectionRefused => TransportError::Io(IoErrorKind::ConnectionRefused),
K::NetworkUnreachable | K::HostUnreachable => {
TransportError::Io(IoErrorKind::NetworkUnreachable)
}
K::WouldBlock => TransportError::Io(IoErrorKind::WouldBlock),
_ => TransportError::Io(IoErrorKind::Other),
};
match kind {
K::TimedOut | K::Interrupted | K::ConnectionRefused => {
crate::log::debug!(
"tokio transport io error: {e} (raw_os={:?}, kind={:?}) mapped to {mapped}",
e.raw_os_error(),
kind,
);
}
_ => {
crate::log::warn!(
"tokio transport io error: {e} (raw_os={:?}, kind={:?}) mapped to {mapped}",
e.raw_os_error(),
kind,
);
}
}
mapped
}
#[derive(Clone, Copy)]
pub struct TokioChannels;
pub struct TokioOneshotReceiver<T>(pub(crate) tokio::sync::oneshot::Receiver<T>);
pub struct TokioUnboundedReceiver<T>(pub(crate) tokio::sync::mpsc::UnboundedReceiver<T>);
impl<T: Send + 'static> OneshotSend<T> for tokio::sync::oneshot::Sender<T> {
fn send(self, value: T) -> Result<(), T> {
tokio::sync::oneshot::Sender::send(self, value)
}
}
impl<T: Send + 'static> OneshotRecv<T> for TokioOneshotReceiver<T> {
async fn recv(self) -> Result<T, OneshotCancelled> {
self.0.await.map_err(|_| OneshotCancelled)
}
}
impl<T: Send + 'static> MpscSend<T> for tokio::sync::mpsc::Sender<T> {
async fn send(&self, value: T) -> Result<(), ()> {
tokio::sync::mpsc::Sender::send(self, value)
.await
.map_err(|_| ())
}
}
impl<T: Send + 'static> MpscRecv<T> for tokio::sync::mpsc::Receiver<T> {
fn recv(&mut self) -> impl Future<Output = Option<T>> + Send + '_ {
self.recv()
}
fn poll_recv(&mut self, cx: &mut core::task::Context<'_>) -> core::task::Poll<Option<T>> {
self.poll_recv(cx)
}
}
impl<T: Send + 'static> UnboundedSend<T> for tokio::sync::mpsc::UnboundedSender<T> {
fn send_now(&self, value: T) -> Result<(), T> {
self.send(value).map_err(|e| e.0)
}
}
impl<T: Send + 'static> UnboundedRecv<T> for TokioUnboundedReceiver<T> {
fn recv(&mut self) -> impl Future<Output = Option<T>> + Send + '_ {
self.0.recv()
}
}
impl ChannelFactory for TokioChannels {
type OneshotSender<T: Send + 'static> = tokio::sync::oneshot::Sender<T>;
type OneshotReceiver<T: Send + 'static> = TokioOneshotReceiver<T>;
type BoundedSender<T: Send + 'static, const N: usize> = tokio::sync::mpsc::Sender<T>;
type BoundedReceiver<T: Send + 'static, const N: usize> = tokio::sync::mpsc::Receiver<T>;
type UnboundedSender<T: Send + 'static> = tokio::sync::mpsc::UnboundedSender<T>;
type UnboundedReceiver<T: Send + 'static> = TokioUnboundedReceiver<T>;
}
impl<T: Send + 'static> crate::transport::OneshotPooled<TokioChannels> for T {
fn oneshot_pair() -> (
<TokioChannels as ChannelFactory>::OneshotSender<T>,
<TokioChannels as ChannelFactory>::OneshotReceiver<T>,
) {
let (tx, rx) = tokio::sync::oneshot::channel();
(tx, TokioOneshotReceiver(rx))
}
}
impl<T: Send + 'static, const N: usize> crate::transport::BoundedPooled<TokioChannels, N> for T {
fn bounded_pair() -> (
<TokioChannels as ChannelFactory>::BoundedSender<T, N>,
<TokioChannels as ChannelFactory>::BoundedReceiver<T, N>,
) {
tokio::sync::mpsc::channel(N)
}
}
impl<T: Send + 'static> crate::transport::UnboundedPooled<TokioChannels> for T {
fn unbounded_pair() -> (
<TokioChannels as ChannelFactory>::UnboundedSender<T>,
<TokioChannels as ChannelFactory>::UnboundedReceiver<T>,
) {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
(tx, TokioUnboundedReceiver(rx))
}
}
use std::sync::Arc;
use crate::buffer_pool::{BufferLease, BufferPool};
use crate::transport::BufferProvider;
#[derive(Clone, Debug)]
pub struct TokioBufferProvider(Arc<BufferPool<10, { crate::UDP_BUFFER_SIZE }>>);
impl TokioBufferProvider {
#[must_use]
pub fn new() -> Self {
Self(Arc::new(BufferPool::new()))
}
}
impl Default for TokioBufferProvider {
fn default() -> Self {
Self::new()
}
}
impl BufferProvider for TokioBufferProvider {
fn claim(&self) -> Option<BufferLease> {
self.0.claim_arc()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn bind_ephemeral_and_report_local_addr() {
let factory = TokioTransport;
let sock = factory
.bind(
SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0),
&SocketOptions::default(),
)
.await
.expect("bind");
let addr = sock.local_addr().expect("local_addr");
assert_eq!(*addr.ip(), Ipv4Addr::LOCALHOST);
assert_ne!(addr.port(), 0, "kernel must assign a non-zero port");
}
#[tokio::test]
async fn round_trip_send_recv_between_two_sockets() {
let factory = TokioTransport;
let opts = SocketOptions::default();
let recv = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0), &opts)
.await
.unwrap();
let recv_addr = recv.local_addr().unwrap();
let send = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0), &opts)
.await
.unwrap();
let payload = b"hello tokio transport";
send.send_to(payload, recv_addr).await.unwrap();
let mut buf = [0u8; 64];
let datagram = tokio::time::timeout(Duration::from_secs(2), recv.recv_from(&mut buf))
.await
.expect("recv timed out")
.expect("recv failed");
assert_eq!(datagram.bytes_received, payload.len());
assert_eq!(&buf[..datagram.bytes_received], payload);
assert!(!datagram.truncated);
}
#[tokio::test]
async fn reuse_address_option_allows_rebind_pattern() {
let opts = SocketOptions {
reuse_address: true,
..SocketOptions::default()
};
let factory = TokioTransport;
let a = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0), &opts)
.await
.unwrap();
let port = a.local_addr().unwrap().port();
let b = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, port), &opts)
.await;
match b {
Ok(_) | Err(TransportError::AddressInUse) => {}
Err(other) => panic!("unexpected rebind error: {other:?}"),
}
drop(a);
}
#[tokio::test]
async fn multicast_loop_v4_option_propagates_in_both_directions() {
let factory = TokioTransport;
let opts_off = SocketOptions {
multicast_loop_v4: Some(false),
multicast_if_v4: Some(Ipv4Addr::LOCALHOST),
..SocketOptions::default()
};
let sock_off = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0), &opts_off)
.await
.expect("bind off");
assert!(
!sock_off.multicast_loop_v4().expect("read off flag"),
"multicast_loop_v4=false must disable IP_MULTICAST_LOOP"
);
let opts_on = SocketOptions {
multicast_loop_v4: Some(true),
multicast_if_v4: Some(Ipv4Addr::LOCALHOST),
..SocketOptions::default()
};
let sock_on = factory
.bind(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0), &opts_on)
.await
.expect("bind on");
assert!(
sock_on.multicast_loop_v4().expect("read on flag"),
"multicast_loop_v4=true must enable IP_MULTICAST_LOOP"
);
}
#[tokio::test]
async fn timer_sleep_elapses_at_least_requested() {
let timer = TokioTimer;
let started = tokio::time::Instant::now();
timer.sleep(Duration::from_millis(25)).await;
assert!(started.elapsed() >= Duration::from_millis(25));
}
#[test]
fn map_io_error_covers_common_kinds() {
use std::io::{Error, ErrorKind};
assert!(matches!(
map_io_error(&Error::from(ErrorKind::AddrInUse)),
TransportError::AddressInUse
));
assert!(matches!(
map_io_error(&Error::from(ErrorKind::TimedOut)),
TransportError::Io(IoErrorKind::TimedOut)
));
assert!(matches!(
map_io_error(&Error::from(ErrorKind::ConnectionRefused)),
TransportError::Io(IoErrorKind::ConnectionRefused)
));
assert!(matches!(
map_io_error(&Error::from(ErrorKind::Unsupported)),
TransportError::Unsupported
));
assert!(matches!(
map_io_error(&Error::from(ErrorKind::Other)),
TransportError::Io(IoErrorKind::Other)
));
}
#[tokio::test]
async fn panic_logging_fut_passes_through_normal_completion() {
use core::future::Future as _;
use core::pin::pin;
use core::task::{Context, Poll};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
let poll_count = Arc::new(AtomicUsize::new(0));
let poll_count_clone = poll_count.clone();
let inner = async move {
poll_count_clone.fetch_add(1, Ordering::SeqCst);
};
let fut = PanicLoggingFut { inner };
let mut fut = pin!(fut);
let waker = futures_util::task::noop_waker();
let mut cx = Context::from_waker(&waker);
match fut.as_mut().poll(&mut cx) {
Poll::Ready(()) => {}
Poll::Pending => panic!(
"PanicLoggingFut wrapping a Ready future returned Pending; \
wrapper is not forwarding `inner.poll` correctly",
),
}
assert_eq!(
poll_count.load(Ordering::SeqCst),
1,
"inner future must have been polled exactly once",
);
}
#[tokio::test]
async fn panic_logging_fut_catches_panic_and_resolves_cleanly() {
use core::future::Future as _;
use core::pin::pin;
use core::task::{Context, Poll};
use std::boxed::Box;
let prev_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let inner = async {
panic!("intentional test panic — must be caught by PanicLoggingFut");
};
let fut = PanicLoggingFut { inner };
let mut fut = pin!(fut);
let waker = futures_util::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let result = fut.as_mut().poll(&mut cx);
std::panic::set_hook(prev_hook);
match result {
Poll::Ready(()) => {}
Poll::Pending => panic!(
"PanicLoggingFut on a panicking future returned Pending; \
expected Ready(()) from the catch_unwind Err arm",
),
}
}
#[tokio::test]
async fn tokio_spawner_isolates_panicking_tasks_from_runtime() {
use crate::transport::Spawner;
use std::boxed::Box;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
let prev_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
TokioSpawner.spawn(async {
panic!("intentional test panic in spawned task");
});
let healthy_done = Arc::new(AtomicBool::new(false));
let healthy_clone = healthy_done.clone();
TokioSpawner.spawn(async move {
healthy_clone.store(true, Ordering::SeqCst);
});
let observed = tokio::time::timeout(Duration::from_secs(1), async {
while !healthy_done.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await;
std::panic::set_hook(prev_hook);
observed.expect(
"healthy task spawned after a panicking one must still complete; \
a hang here means the panic took down the runtime — \
PanicLoggingFut wrapper missing or broken",
);
}
}