diff --git a/src/server/rwhod/packet_sender.rs b/src/server/rwhod/packet_sender.rs index 4ed9480..da4c9b2 100644 --- a/src/server/rwhod/packet_sender.rs +++ b/src/server/rwhod/packet_sender.rs @@ -1,7 +1,7 @@ use nix::{ifaddrs::getifaddrs, net::if_::InterfaceFlags}; use std::{ collections::HashSet, - net::{IpAddr, SocketAddr}, + net::{IpAddr, Ipv4Addr, SocketAddr}, sync::Arc, }; use tokio::{ @@ -29,6 +29,11 @@ pub struct RwhodSendTarget { pub addr: IpAddr, } +/// Computes the broadcast address for an IPv4 address/netmask pair +fn ipv4_broadcast_address(address: Ipv4Addr, netmask: Ipv4Addr) -> Ipv4Addr { + Ipv4Addr::from(u32::from(address) | !u32::from(netmask)) +} + /// Find all networks network interfaces suitable for rwhod communication. /// /// If `allowed_interfaces` is `Some`, only interfaces whose name is contained @@ -53,24 +58,29 @@ pub fn determine_relevant_interfaces( }) .filter_map(|iface| { let neighbor_addr = if iface.flags.contains(InterfaceFlags::IFF_BROADCAST) { - iface.broadcast + match ( + iface.address.as_ref().and_then(|a| a.as_sockaddr_in()), + iface.netmask.as_ref().and_then(|a| a.as_sockaddr_in()), + ) { + (Some(addr), Some(mask)) => { + Some(ipv4_broadcast_address(addr.ip(), mask.ip()).into()) + } + _ => None, + } } else if iface.flags.contains(InterfaceFlags::IFF_POINTOPOINT) { - iface.destination + iface.destination.and_then(|addr| { + addr.as_sockaddr_in() + .map(|sa| IpAddr::V4(sa.ip())) + .or_else(|| addr.as_sockaddr_in6().map(|sa| IpAddr::V6(sa.ip()))) + }) } else { None }; - match neighbor_addr { - Some(addr) => addr - .as_sockaddr_in() - .map(|sa| IpAddr::V4(sa.ip())) - .or_else(|| addr.as_sockaddr_in6().map(|sa| IpAddr::V6(sa.ip()))) - .map(|ip_addr| RwhodSendTarget { - name: iface.interface_name, - addr: ip_addr, - }), - None => None, - } + neighbor_addr.map(|ip_addr| RwhodSendTarget { + name: iface.interface_name, + addr: ip_addr, + }) }) // keep first occurrence per interface name .scan(HashSet::new(), |seen, n| { @@ -167,6 +177,30 @@ pub async fn rwhod_packet_sender_task( mod tests { use super::*; + #[test] + fn test_ipv4_broadcast_address() { + assert_eq!( + ipv4_broadcast_address( + Ipv4Addr::new(192, 168, 1, 2), + Ipv4Addr::new(255, 255, 255, 0) + ), + Ipv4Addr::new(192, 168, 1, 255) + ); + + assert_eq!( + ipv4_broadcast_address(Ipv4Addr::new(10, 0, 5, 200), Ipv4Addr::new(255, 0, 0, 0)), + Ipv4Addr::new(10, 255, 255, 255) + ); + + assert_eq!( + ipv4_broadcast_address( + Ipv4Addr::new(203, 0, 113, 42), + Ipv4Addr::new(255, 255, 255, 255) + ), + Ipv4Addr::new(203, 0, 113, 42) + ); + } + #[test] fn test_determine_relevant_interfaces() { let interfaces = determine_relevant_interfaces(None).unwrap();