use std::net::SocketAddr;
use rtime_core::timestamp::NtpTimestamp;
use tokio::net::UdpSocket;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TimestampMode {
Userspace,
Software,
Hardware,
}
pub struct TimestampedSocket {
inner: UdpSocket,
mode: TimestampMode,
}
pub struct ReceivedPacket {
pub data: Vec<u8>,
pub addr: SocketAddr,
pub recv_time: NtpTimestamp,
}
impl TimestampedSocket {
pub async fn bind(addr: &str) -> Result<Self, std::io::Error> {
let socket = UdpSocket::bind(addr).await?;
Ok(Self {
inner: socket,
mode: TimestampMode::Userspace,
})
}
pub fn from_socket(socket: UdpSocket) -> Self {
Self {
inner: socket,
mode: TimestampMode::Userspace,
}
}
pub async fn send_to(
&self,
buf: &[u8],
addr: SocketAddr,
) -> Result<NtpTimestamp, std::io::Error> {
self.inner.send_to(buf, addr).await?;
let ts = NtpTimestamp::now();
Ok(ts)
}
pub async fn recv_from(&self, buf: &mut [u8]) -> Result<ReceivedPacket, std::io::Error> {
let (len, addr) = self.inner.recv_from(buf).await?;
let recv_time = NtpTimestamp::now();
Ok(ReceivedPacket {
data: buf[..len].to_vec(),
addr,
recv_time,
})
}
pub fn local_addr(&self) -> Result<SocketAddr, std::io::Error> {
self.inner.local_addr()
}
pub fn inner(&self) -> &UdpSocket {
&self.inner
}
pub fn timestamp_mode(&self) -> TimestampMode {
self.mode
}
#[cfg(target_os = "linux")]
pub fn enable_software_timestamps(&mut self) -> std::io::Result<()> {
nix::sys::socket::setsockopt(
&self.inner,
nix::sys::socket::sockopt::ReceiveTimestampns,
&true,
)
.map_err(std::io::Error::from)?;
self.mode = TimestampMode::Software;
Ok(())
}
#[cfg(target_os = "freebsd")]
pub fn enable_software_timestamps(&mut self) -> std::io::Result<()> {
nix::sys::socket::setsockopt(
&self.inner,
nix::sys::socket::sockopt::ReceiveTimestamp,
&true,
)
.map_err(std::io::Error::from)?;
self.mode = TimestampMode::Software;
Ok(())
}
#[cfg(not(any(target_os = "linux", target_os = "freebsd")))]
pub fn enable_software_timestamps(&mut self) -> std::io::Result<()> {
Ok(())
}
#[cfg(target_os = "linux")]
pub fn enable_hardware_timestamps(&mut self) -> std::io::Result<bool> {
use nix::sys::socket::{TimestampingFlag, setsockopt, sockopt};
let hw_flags = TimestampingFlag::SOF_TIMESTAMPING_TX_HARDWARE
| TimestampingFlag::SOF_TIMESTAMPING_RX_HARDWARE
| TimestampingFlag::SOF_TIMESTAMPING_RAW_HARDWARE;
if setsockopt(&self.inner, sockopt::Timestamping, &hw_flags).is_ok() {
self.mode = TimestampMode::Hardware;
return Ok(true);
}
let sw_flags = TimestampingFlag::SOF_TIMESTAMPING_TX_SOFTWARE
| TimestampingFlag::SOF_TIMESTAMPING_RX_SOFTWARE
| TimestampingFlag::SOF_TIMESTAMPING_SOFTWARE;
setsockopt(&self.inner, sockopt::Timestamping, &sw_flags).map_err(std::io::Error::from)?;
self.mode = TimestampMode::Software;
Ok(false)
}
#[cfg(target_os = "freebsd")]
pub fn enable_hardware_timestamps(&mut self) -> std::io::Result<bool> {
self.enable_software_timestamps()?;
Ok(false)
}
#[cfg(not(any(target_os = "linux", target_os = "freebsd")))]
pub fn enable_hardware_timestamps(&mut self) -> std::io::Result<bool> {
Ok(false)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn bind_ephemeral_port() {
let sock = TimestampedSocket::bind("127.0.0.1:0")
.await
.expect("bind should succeed");
let addr = sock.local_addr().expect("local_addr should succeed");
assert!(addr.port() > 0, "should have been assigned a port");
assert_eq!(addr.ip(), std::net::Ipv4Addr::LOCALHOST);
}
#[tokio::test]
async fn send_recv_loopback() {
let sender = TimestampedSocket::bind("127.0.0.1:0")
.await
.expect("bind sender");
let receiver = TimestampedSocket::bind("127.0.0.1:0")
.await
.expect("bind receiver");
let receiver_addr = receiver.local_addr().expect("receiver local_addr");
let payload = b"NTP test packet data";
let send_ts = sender
.send_to(payload, receiver_addr)
.await
.expect("send_to");
let mut buf = [0u8; 1024];
let received = receiver.recv_from(&mut buf).await.expect("recv_from");
assert_eq!(received.data, payload);
assert_eq!(received.addr, sender.local_addr().expect("sender addr"));
assert_ne!(send_ts, NtpTimestamp::ZERO);
assert_ne!(received.recv_time, NtpTimestamp::ZERO);
assert!(
received.recv_time.raw() >= send_ts.raw(),
"recv_time ({:?}) should be >= send_ts ({:?})",
received.recv_time,
send_ts
);
}
#[tokio::test]
async fn from_socket_works() {
let raw = UdpSocket::bind("127.0.0.1:0")
.await
.expect("bind raw socket");
let addr = raw.local_addr().expect("raw local_addr");
let wrapped = TimestampedSocket::from_socket(raw);
assert_eq!(wrapped.local_addr().expect("wrapped local_addr"), addr);
}
#[tokio::test]
async fn multiple_packets_loopback() {
let sender = TimestampedSocket::bind("127.0.0.1:0")
.await
.expect("bind sender");
let receiver = TimestampedSocket::bind("127.0.0.1:0")
.await
.expect("bind receiver");
let receiver_addr = receiver.local_addr().expect("receiver addr");
for i in 0u8..5 {
let payload = [i; 48]; sender.send_to(&payload, receiver_addr).await.expect("send");
let mut buf = [0u8; 1024];
let pkt = receiver.recv_from(&mut buf).await.expect("recv");
assert_eq!(pkt.data.len(), 48);
assert_eq!(pkt.data[0], i);
}
}
}