use std::net::Ipv4Addr;
use std::os::unix::io::AsRawFd;
use nix::sys::socket::{IpMembershipRequest, setsockopt, sockopt};
use tokio::net::UdpSocket;
pub const PTP_PRIMARY_MULTICAST_V4: Ipv4Addr = Ipv4Addr::new(224, 0, 1, 129);
pub const PTP_PDELAY_MULTICAST_V4: Ipv4Addr = Ipv4Addr::new(224, 0, 0, 107);
pub const PTP_EVENT_PORT: u16 = 319;
pub const PTP_GENERAL_PORT: u16 = 320;
pub fn join_multicast(
socket: &UdpSocket,
multicast_addr: Ipv4Addr,
interface: Ipv4Addr,
) -> std::io::Result<()> {
let interface_opt = if interface.is_unspecified() {
None
} else {
Some(interface)
};
let mreq = IpMembershipRequest::new(multicast_addr, interface_opt);
setsockopt(socket, sockopt::IpAddMembership, &mreq).map_err(std::io::Error::from)
}
pub fn leave_multicast(
socket: &UdpSocket,
multicast_addr: Ipv4Addr,
interface: Ipv4Addr,
) -> std::io::Result<()> {
let interface_opt = if interface.is_unspecified() {
None
} else {
Some(interface)
};
let mreq = IpMembershipRequest::new(multicast_addr, interface_opt);
setsockopt(socket, sockopt::IpDropMembership, &mreq).map_err(std::io::Error::from)
}
pub fn set_multicast_interface(socket: &UdpSocket, interface: Ipv4Addr) -> std::io::Result<()> {
let addr = libc::in_addr {
s_addr: u32::from_ne_bytes(interface.octets()),
};
let fd = socket.as_raw_fd();
let ret = unsafe {
libc::setsockopt(
fd,
libc::IPPROTO_IP,
libc::IP_MULTICAST_IF,
&addr as *const libc::in_addr as *const libc::c_void,
std::mem::size_of::<libc::in_addr>() as libc::socklen_t,
)
};
if ret < 0 {
return Err(std::io::Error::last_os_error());
}
Ok(())
}
pub fn set_multicast_loopback(socket: &UdpSocket, enabled: bool) -> std::io::Result<()> {
setsockopt(socket, sockopt::IpMulticastLoop, &enabled).map_err(std::io::Error::from)
}
pub fn set_multicast_ttl(socket: &UdpSocket, ttl: u8) -> std::io::Result<()> {
setsockopt(socket, sockopt::IpMulticastTtl, &ttl).map_err(std::io::Error::from)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn constants_are_correct() {
assert_eq!(PTP_PRIMARY_MULTICAST_V4, Ipv4Addr::new(224, 0, 1, 129));
assert_eq!(PTP_PDELAY_MULTICAST_V4, Ipv4Addr::new(224, 0, 0, 107));
assert_eq!(PTP_EVENT_PORT, 319);
assert_eq!(PTP_GENERAL_PORT, 320);
}
#[tokio::test]
async fn join_and_leave_loopback() {
let socket = UdpSocket::bind("0.0.0.0:0")
.await
.expect("bind should succeed");
let result = join_multicast(&socket, PTP_PRIMARY_MULTICAST_V4, Ipv4Addr::LOCALHOST);
if let Err(ref e) = result {
eprintln!("join_multicast failed (expected in some CI): {e}");
return;
}
result.unwrap();
leave_multicast(&socket, PTP_PRIMARY_MULTICAST_V4, Ipv4Addr::LOCALHOST)
.expect("leave should succeed after join");
}
#[tokio::test]
async fn set_loopback_and_ttl() {
let socket = UdpSocket::bind("0.0.0.0:0")
.await
.expect("bind should succeed");
set_multicast_loopback(&socket, false).expect("disable loopback");
set_multicast_loopback(&socket, true).expect("enable loopback");
set_multicast_ttl(&socket, 1).expect("set ttl=1");
set_multicast_ttl(&socket, 128).expect("set ttl=128");
}
}