use std::io;
use std::net::{Ipv4Addr, SocketAddr};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SearchDatagram {
pub n: usize,
pub src: SocketAddr,
pub iface_ip: Option<Ipv4Addr>,
pub drops: u32,
}
pub struct SearchUdpSocket(sys::Sock);
impl SearchUdpSocket {
pub fn bind_ephemeral(
broadcast: bool,
pump_name: &str,
pump_priority: crate::runtime::task::ThreadPriority,
) -> io::Result<Self> {
sys::Sock::bind_ephemeral(broadcast, pump_name, pump_priority).map(Self)
}
pub fn bind_beacon(
port: u16,
pump_name: &str,
pump_priority: crate::runtime::task::ThreadPriority,
) -> io::Result<Self> {
sys::Sock::bind_beacon(port, pump_name, pump_priority).map(Self)
}
pub fn join_multicast_v4(&self, group: Ipv4Addr, iface: Ipv4Addr) -> io::Result<()> {
self.0.join_multicast_v4(group, iface)
}
pub fn set_recv_buffer_size(&self, size: usize) -> io::Result<()> {
self.0.set_recv_buffer_size(size)
}
pub fn set_multicast_ttl_v4(&self, ttl: u32) -> io::Result<()> {
self.0.set_multicast_ttl_v4(ttl)
}
pub fn enable_so_rxq_ovfl(&self) -> io::Result<()> {
self.0.enable_so_rxq_ovfl()
}
pub fn local_addrs(&self) -> Vec<SocketAddr> {
self.0.local_addrs()
}
pub async fn recv(&self, buf: &mut [u8]) -> io::Result<SearchDatagram> {
self.0.recv(buf).await
}
pub async fn send_to(&self, buf: &[u8], dest: SocketAddr) -> io::Result<usize> {
self.0.send_to(buf, dest).await
}
pub async fn fanout_to(
&self,
buf: &[u8],
dest: SocketAddr,
ifaces: &[Ipv4Addr],
) -> io::Result<usize> {
self.0.fanout_to(buf, dest, ifaces).await
}
}
#[cfg(tokio_backend)]
mod sys {
use super::{SearchDatagram, SocketAddr, io};
use crate::net::async_udp_v4::AsyncUdpV4;
pub(super) struct Sock(AsyncUdpV4);
impl Sock {
pub(super) fn bind_ephemeral(
broadcast: bool,
pump_name: &str,
pump_priority: crate::runtime::task::ThreadPriority,
) -> io::Result<Self> {
let _ = (pump_name, pump_priority);
AsyncUdpV4::bind_ephemeral_same_port(broadcast).map(Self)
}
pub(super) fn bind_beacon(
port: u16,
pump_name: &str,
pump_priority: crate::runtime::task::ThreadPriority,
) -> io::Result<Self> {
let _ = (pump_name, pump_priority);
AsyncUdpV4::bind_non_loopback(port, true).map(Self)
}
pub(super) fn join_multicast_v4(
&self,
group: std::net::Ipv4Addr,
iface: std::net::Ipv4Addr,
) -> io::Result<()> {
if iface.is_unspecified() {
self.0.join_multicast_v4(group)
} else {
self.0.join_multicast_v4_on(group, iface)
}
}
pub(super) fn set_recv_buffer_size(&self, size: usize) -> io::Result<()> {
self.0.set_recv_buffer_size(size)
}
pub(super) fn set_multicast_ttl_v4(&self, ttl: u32) -> io::Result<()> {
self.0.set_multicast_ttl_v4(ttl)
}
pub(super) fn enable_so_rxq_ovfl(&self) -> io::Result<()> {
self.0.enable_so_rxq_ovfl()
}
pub(super) fn local_addrs(&self) -> Vec<SocketAddr> {
self.0.local_addrs()
}
pub(super) async fn recv(&self, buf: &mut [u8]) -> io::Result<SearchDatagram> {
let (meta, drops) = self.0.recv_with_meta_with_drops(buf).await?;
Ok(SearchDatagram {
n: meta.n,
src: meta.src,
iface_ip: Some(meta.iface_ip),
drops,
})
}
pub(super) async fn send_to(&self, buf: &[u8], dest: SocketAddr) -> io::Result<usize> {
self.0.send_to(buf, dest).await
}
pub(super) async fn fanout_to(
&self,
buf: &[u8],
dest: SocketAddr,
ifaces: &[std::net::Ipv4Addr],
) -> io::Result<usize> {
if ifaces.is_empty() {
return self.0.fanout_to(buf, dest).await;
}
let mut ok = 0usize;
let mut last_err: Option<io::Error> = None;
for ip in ifaces {
if ip.is_loopback() {
continue;
}
match self.0.send_via(buf, dest, *ip).await {
Ok(_) => ok += 1,
Err(e) => last_err = Some(e),
}
}
if ok == 0 {
return Err(last_err.unwrap_or_else(|| {
io::Error::new(
io::ErrorKind::AddrNotAvailable,
"SEARCH fanout: no listed interface available",
)
}));
}
Ok(ok)
}
}
}
#[cfg(exec_backend)]
mod sys {
use super::{Ipv4Addr, SearchDatagram, SocketAddr, io};
use crate::runtime::task::{StackSizeClass, ThreadPriority, spawn_dedicated_thread};
use std::net::UdpSocket;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
const PUMP_WAKE_INTERVAL: Duration = Duration::from_millis(200);
const RECV_BUF: usize = 64 * 1024;
pub(super) struct Sock {
socket: Arc<UdpSocket>,
rx: tokio::sync::Mutex<tokio::sync::mpsc::UnboundedReceiver<(Vec<u8>, SocketAddr)>>,
stop: Arc<AtomicBool>,
pump: Option<std::thread::JoinHandle<()>>,
}
impl Sock {
pub(super) fn bind_ephemeral(
broadcast: bool,
pump_name: &str,
pump_priority: ThreadPriority,
) -> io::Result<Self> {
let socket = UdpSocket::bind((Ipv4Addr::UNSPECIFIED, 0))?;
if broadcast {
socket.set_broadcast(true)?;
}
Self::with_pump(socket, pump_name, pump_priority)
}
pub(super) fn bind_beacon(
port: u16,
pump_name: &str,
pump_priority: ThreadPriority,
) -> io::Result<Self> {
let socket = bind_shared_port(port)?;
socket.set_broadcast(true)?;
Self::with_pump(socket, pump_name, pump_priority)
}
pub(super) fn join_multicast_v4(&self, group: Ipv4Addr, iface: Ipv4Addr) -> io::Result<()> {
self.socket.join_multicast_v4(&group, &iface)
}
fn with_pump(
socket: UdpSocket,
pump_name: &str,
pump_priority: ThreadPriority,
) -> io::Result<Self> {
socket.set_read_timeout(Some(PUMP_WAKE_INTERVAL))?;
let socket = Arc::new(socket);
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
let stop = Arc::new(AtomicBool::new(false));
let pump_socket = Arc::clone(&socket);
let pump_stop = Arc::clone(&stop);
let pump = spawn_dedicated_thread(
pump_name.to_string(),
pump_priority,
StackSizeClass::Medium,
move || pump(&pump_socket, &pump_stop, &tx),
)?;
Ok(Self {
socket,
rx: tokio::sync::Mutex::new(rx),
stop,
pump: Some(pump),
})
}
pub(super) fn set_recv_buffer_size(&self, size: usize) -> io::Result<()> {
set_int_opt(
&self.socket,
sockopt::SOL_SOCKET,
sockopt::SO_RCVBUF,
size as _,
)
}
pub(super) fn set_multicast_ttl_v4(&self, ttl: u32) -> io::Result<()> {
set_int_opt(
&self.socket,
sockopt::IPPROTO_IP,
sockopt::IP_MULTICAST_TTL,
ttl as _,
)
}
pub(super) fn enable_so_rxq_ovfl(&self) -> io::Result<()> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"SO_RXQ_OVFL needs a recvmsg receive path; this SEARCH socket uses recv_from",
))
}
pub(super) fn local_addrs(&self) -> Vec<SocketAddr> {
self.socket.local_addr().into_iter().collect()
}
pub(super) async fn recv(&self, buf: &mut [u8]) -> io::Result<SearchDatagram> {
let Some((bytes, src)) = self.rx.lock().await.recv().await else {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"SEARCH receive pump stopped",
));
};
let n = bytes.len().min(buf.len());
buf[..n].copy_from_slice(&bytes[..n]);
Ok(SearchDatagram {
n,
src,
iface_ip: None,
drops: 0,
})
}
pub(super) async fn send_to(&self, buf: &[u8], dest: SocketAddr) -> io::Result<usize> {
self.socket.send_to(buf, dest)
}
pub(super) async fn fanout_to(
&self,
buf: &[u8],
dest: SocketAddr,
ifaces: &[Ipv4Addr],
) -> io::Result<usize> {
let port = dest.port();
let dests: Vec<SocketAddr> = match dest {
SocketAddr::V4(v4) if v4.ip().is_broadcast() => eligible_broadcast_addrs(ifaces)
.into_iter()
.map(|ip| SocketAddr::from((ip, port)))
.collect(),
_ => vec![dest],
};
let mut ok = 0usize;
let mut last_err: Option<io::Error> = None;
for d in dests {
match self.socket.send_to(buf, d) {
Ok(_) => ok += 1,
Err(e) => {
tracing::debug!(
target: "epics_base_rs::net",
dest = %d,
error = %e,
"SEARCH fanout send failed"
);
last_err = Some(e);
}
}
}
if ok == 0 {
return Err(last_err.unwrap_or_else(|| {
io::Error::new(
io::ErrorKind::AddrNotAvailable,
"SEARCH fanout: no eligible broadcast destination",
)
}));
}
Ok(ok)
}
}
impl Drop for Sock {
fn drop(&mut self) {
self.stop.store(true, Ordering::Release);
if let Some(pump) = self.pump.take() {
let _ = pump.join();
}
}
}
fn eligible_broadcast_addrs(ifaces: &[Ipv4Addr]) -> Vec<Ipv4Addr> {
if ifaces.is_empty() {
return crate::net::iface_v4::broadcast_addrs();
}
let Ok(all) = crate::net::iface_v4::enumerate() else {
return Vec::new();
};
let mut out: Vec<Ipv4Addr> = Vec::new();
for iface in all {
if !ifaces.contains(&iface.ip) {
continue;
}
if let Some(dest) = iface.search_destination() {
if !out.contains(&dest) {
out.push(dest);
}
}
}
out
}
fn pump(
socket: &UdpSocket,
stop: &AtomicBool,
tx: &tokio::sync::mpsc::UnboundedSender<(Vec<u8>, SocketAddr)>,
) {
let mut buf = vec![0u8; RECV_BUF];
while !stop.load(Ordering::Acquire) {
match socket.recv_from(&mut buf) {
Ok((n, src)) => {
if tx.send((buf[..n].to_vec(), src)).is_err() {
return;
}
}
Err(e) if is_wake_timeout(e.kind()) => continue,
Err(e)
if matches!(
e.kind(),
io::ErrorKind::ConnectionRefused
| io::ErrorKind::ConnectionReset
| io::ErrorKind::Interrupted
) =>
{
continue;
}
Err(e) => {
tracing::warn!(
target: "epics_base_rs::net",
error = %e,
"SEARCH receive pump stopping"
);
return;
}
}
}
}
fn is_wake_timeout(kind: io::ErrorKind) -> bool {
matches!(kind, io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut)
}
#[cfg(unix)]
mod sockopt {
pub(super) const SOL_SOCKET: libc::c_int = libc::SOL_SOCKET;
pub(super) const SO_RCVBUF: libc::c_int = libc::SO_RCVBUF;
pub(super) const IPPROTO_IP: libc::c_int = libc::IPPROTO_IP;
pub(super) const IP_MULTICAST_TTL: libc::c_int = libc::IP_MULTICAST_TTL;
}
#[cfg(unix)]
fn bind_shared_port(port: u16) -> io::Result<UdpSocket> {
use std::os::fd::FromRawFd;
let fd = unsafe { libc::socket(libc::AF_INET, libc::SOCK_DGRAM, libc::IPPROTO_UDP) };
if fd < 0 {
return Err(io::Error::last_os_error());
}
let socket = unsafe { UdpSocket::from_raw_fd(fd) };
set_int_opt(&socket, libc::SOL_SOCKET, libc::SO_REUSEADDR, 1)?;
set_int_opt(&socket, libc::SOL_SOCKET, libc::SO_REUSEPORT, 1)?;
let mut sin: libc::sockaddr_in = unsafe { std::mem::zeroed() };
sin.sin_family = libc::AF_INET as libc::sa_family_t;
sin.sin_port = port.to_be();
let rc = unsafe {
libc::bind(
fd,
std::ptr::addr_of!(sin).cast(),
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(io::Error::last_os_error());
}
Ok(socket)
}
#[cfg(not(unix))]
fn bind_shared_port(_port: u16) -> io::Result<UdpSocket> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"port-sharing bind on a non-Unix exec-backend socket",
))
}
#[cfg(unix)]
fn set_int_opt(
socket: &UdpSocket,
level: libc::c_int,
name: libc::c_int,
value: libc::c_int,
) -> io::Result<()> {
use std::os::fd::AsRawFd;
let rc = unsafe {
libc::setsockopt(
socket.as_raw_fd(),
level,
name,
std::ptr::addr_of!(value).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
#[cfg(not(unix))]
mod sockopt {
pub(super) const SOL_SOCKET: i32 = 0;
pub(super) const SO_RCVBUF: i32 = 0;
pub(super) const IPPROTO_IP: i32 = 0;
pub(super) const IP_MULTICAST_TTL: i32 = 0;
}
#[cfg(not(unix))]
fn set_int_opt(_socket: &UdpSocket, _level: i32, _name: i32, _value: i32) -> io::Result<()> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"socket options on a non-Unix exec-backend SEARCH socket",
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::task::ThreadPriority;
fn bind() -> SearchUdpSocket {
SearchUdpSocket::bind_ephemeral(true, "test-CAC-UDP", ThreadPriority::Medium)
.expect("bind ephemeral SEARCH socket")
}
#[epics_macros_rs::epics_test]
async fn binds_and_reports_a_local_address() {
let sock = bind();
let addrs = sock.local_addrs();
assert!(!addrs.is_empty(), "a bound SEARCH socket has an address");
for a in &addrs {
assert_ne!(a.port(), 0, "ephemeral bind assigns a real port: {a}");
}
}
#[epics_macros_rs::epics_test]
async fn beacon_binds_the_port_it_was_given() {
let port = {
let probe = bind();
probe
.local_addrs()
.iter()
.find(|a| a.is_ipv4())
.expect("an IPv4 SEARCH address")
.port()
};
let beacon = SearchUdpSocket::bind_beacon(port, "test-BEACON", ThreadPriority::Medium)
.expect("beacon bind");
assert!(
beacon.local_addrs().iter().any(|a| a.port() == port),
"a beacon listener must bind the port it was given; got {:?}",
beacon.local_addrs()
);
}
#[epics_macros_rs::epics_test]
async fn round_trips_a_datagram_to_itself() {
let sock = bind();
let port = sock
.local_addrs()
.iter()
.find(|a| a.is_ipv4())
.expect("an IPv4 SEARCH address")
.port();
let dest = SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, port));
sock.send_to(b"CA-SEARCH", dest).await.expect("send_to");
let mut buf = [0u8; 64];
let dg =
crate::runtime::task::timeout(std::time::Duration::from_secs(5), sock.recv(&mut buf))
.await
.expect("a datagram sent to ourselves arrives")
.expect("recv");
assert_eq!(&buf[..dg.n], b"CA-SEARCH");
assert_eq!(dg.drops, 0, "a single quiet datagram overflows nothing");
}
}