use bytesize::ByteSize;
use pnet::datalink::{self, NetworkInterface};
use std::cmp::min;
use std::net::IpAddr;
use std::time::Duration;
use sysinfo::Networks;
use tokio::sync::watch;
use tracing::{debug, info, warn};
#[derive(Debug, Clone, Default)]
pub struct NetworkStats {
pub max_rx_bandwidth: u64,
pub rx_bandwidth: Option<u64>,
pub max_tx_bandwidth: u64,
pub tx_bandwidth: Option<u64>,
}
#[derive(Debug, Clone)]
pub struct Network {
stats: watch::Receiver<Option<NetworkStats>>,
}
impl Network {
pub fn new(ip: IpAddr, rate_limit: ByteSize) -> Network {
let rate_limit = Self::byte_size_to_bits(rate_limit); let Some(interface) = Self::get_network_interface_by_ip(ip) else {
warn!(
"can not find interface for IP address {}, network interface unknown with bandwidth {} bps",
ip, rate_limit
);
return Self::spawn("unknown".to_string(), rate_limit);
};
match Self::get_speed(&interface.name) {
Some(speed) => {
let bandwidth = min(Self::megabits_to_bits(speed), rate_limit);
info!(
"network interface {} with bandwidth {} bps",
interface.name, bandwidth
);
Self::spawn(interface.name, bandwidth)
}
None => {
warn!(
"can not get speed, network interface {} with bandwidth {} bps",
interface.name, rate_limit
);
Self::spawn(interface.name, rate_limit)
}
}
}
fn spawn(interface_name: String, bandwidth: u64) -> Network {
let (tx, rx) = watch::channel(None);
tokio::spawn(
StatsCollector {
interface_name,
bandwidth,
}
.run(tx),
);
Self { stats: rx }
}
pub async fn get_stats(&mut self) -> Option<NetworkStats> {
self.stats.changed().await.ok()?;
self.stats.borrow_and_update().clone()
}
pub fn get_speed(name: &str) -> Option<u64> {
#[cfg(target_os = "linux")]
{
let speed_path = format!("/sys/class/net/{name}/speed");
std::fs::read_to_string(&speed_path)
.ok()
.and_then(|speed_str| speed_str.trim().parse::<u64>().ok())
}
#[cfg(not(target_os = "linux"))]
{
warn!("can not get interface {} speed on non-linux platform", name);
None
}
}
pub fn get_network_interface_by_ip(ip: IpAddr) -> Option<NetworkInterface> {
datalink::interfaces()
.into_iter()
.find(|interface| interface.ips.iter().any(|ip_net| ip_net.ip() == ip))
}
pub fn byte_size_to_bits(size: ByteSize) -> u64 {
size.as_u64() * 8
}
pub fn megabits_to_bits(size: u64) -> u64 {
size * 1_000_000 }
pub fn bytes_to_bits(size: u64) -> u64 {
size * 8 }
}
#[derive(Debug)]
struct StatsCollector {
interface_name: String,
bandwidth: u64,
}
impl StatsCollector {
const DEFAULT_NETWORK_REFRESH_INTERVAL: Duration = Duration::from_secs(1);
async fn run(self, tx: watch::Sender<Option<NetworkStats>>) {
loop {
let stats = self.collect().await;
if tx.send(Some(stats)).is_err() {
return;
}
}
}
async fn collect(&self) -> NetworkStats {
let mut networks = Networks::new_with_refreshed_list();
tokio::time::sleep(Self::DEFAULT_NETWORK_REFRESH_INTERVAL).await;
networks.refresh(true);
let Some(network_stats) = networks.get(self.interface_name.as_str()) else {
warn!(
"can not find network data for interface {}",
self.interface_name
);
return NetworkStats {
max_rx_bandwidth: self.bandwidth,
max_tx_bandwidth: self.bandwidth,
..Default::default()
};
};
let rx_bandwidth = (Network::bytes_to_bits(network_stats.received()) as f64
/ Self::DEFAULT_NETWORK_REFRESH_INTERVAL.as_secs_f64())
.round() as u64;
let tx_bandwidth = (Network::bytes_to_bits(network_stats.transmitted()) as f64
/ Self::DEFAULT_NETWORK_REFRESH_INTERVAL.as_secs_f64())
.round() as u64;
debug!(
"network interface {} max receive bandwidth: {} bps, receive bandwidth: {} bps, max transmit bandwidth: {} bps, transmit bandwidth: {} bps",
self.interface_name, self.bandwidth, rx_bandwidth, self.bandwidth, tx_bandwidth
);
NetworkStats {
max_rx_bandwidth: self.bandwidth,
rx_bandwidth: Some(rx_bandwidth),
max_tx_bandwidth: self.bandwidth,
tx_bandwidth: Some(tx_bandwidth),
}
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::type_complexity)]
use super::*;
use bytesize::ByteSize;
use std::net::Ipv4Addr;
use std::time::Instant;
use tokio::task::JoinSet;
#[tokio::test]
async fn get_stats_shares_one_collector_across_clones() {
let mut network = Network::new(IpAddr::V4(Ipv4Addr::LOCALHOST), ByteSize::mb(100));
let start = Instant::now();
let mut join_set = JoinSet::new();
for _ in 0..10 {
let mut network = network.clone();
join_set.spawn(async move { network.get_stats().await });
}
while let Some(stats) = join_set.join_next().await {
assert!(stats.unwrap().is_some());
}
assert!(start.elapsed() < StatsCollector::DEFAULT_NETWORK_REFRESH_INTERVAL * 2);
let start = Instant::now();
assert!(network.get_stats().await.is_some());
assert!(start.elapsed() < StatsCollector::DEFAULT_NETWORK_REFRESH_INTERVAL / 2);
let start = Instant::now();
assert!(network.get_stats().await.is_some());
assert!(start.elapsed() >= StatsCollector::DEFAULT_NETWORK_REFRESH_INTERVAL / 2);
assert!(start.elapsed() < StatsCollector::DEFAULT_NETWORK_REFRESH_INTERVAL * 2);
}
#[tokio::test]
async fn new_falls_back_to_the_rate_limit_for_an_unknown_interface() {
let mut network =
Network::new(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)), ByteSize::mb(100));
let stats = network.get_stats().await.unwrap();
assert_eq!(stats.max_rx_bandwidth, 800_000_000);
assert_eq!(stats.max_tx_bandwidth, 800_000_000);
assert_eq!(stats.rx_bandwidth, None);
assert_eq!(stats.tx_bandwidth, None);
}
#[test]
fn get_network_interface_by_ip_finds_only_bound_addresses() {
let test_cases: Vec<(IpAddr, fn(Option<NetworkInterface>))> = vec![
(IpAddr::V4(Ipv4Addr::LOCALHOST), |interface| {
assert!(interface.unwrap().is_loopback());
}),
(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)), |interface| {
assert!(interface.is_none());
}),
];
for (ip, expect) in test_cases {
expect(Network::get_network_interface_by_ip(ip));
}
}
#[test]
fn byte_size_to_bits_multiplies_by_eight() {
let test_cases = vec![
(ByteSize::kb(1), 8_000),
(ByteSize::mb(1), 8_000_000),
(ByteSize::gb(1), 8_000_000_000),
(ByteSize::b(0), 0),
];
for (size, expected) in test_cases {
assert_eq!(Network::byte_size_to_bits(size), expected);
}
}
#[test]
fn megabits_to_bits_multiplies_by_a_million() {
let test_cases = vec![(1, 1_000_000), (1000, 1_000_000_000), (0, 0)];
for (size, expected) in test_cases {
assert_eq!(Network::megabits_to_bits(size), expected);
}
}
#[test]
fn bytes_to_bits_multiplies_by_eight() {
let test_cases = vec![(1, 8), (1000, 8_000), (0, 0)];
for (size, expected) in test_cases {
assert_eq!(Network::bytes_to_bits(size), expected);
}
}
}