use anyhow::Result;
use futures::StreamExt;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Instant;
use super::{SpeedProgress, SpeedResult, SpeedSample};
use crate::utils::format_speed_mbps;
const SAMPLE_INTERVAL_MS: u64 = 250;
pub async fn measure_download<F>(
url: &str,
duration_secs: u64,
mut progress_callback: F,
) -> Result<SpeedResult>
where
F: FnMut(SpeedProgress),
{
use reqwest::header::{HeaderMap, HeaderValue, ACCEPT, ORIGIN, REFERER, USER_AGENT};
let mut headers = HeaderMap::new();
headers.insert(USER_AGENT, HeaderValue::from_static("Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"));
headers.insert(ACCEPT, HeaderValue::from_static("*/*"));
if url.contains("cloudflare") {
headers.insert(
ORIGIN,
HeaderValue::from_static("https://speed.cloudflare.com"),
);
headers.insert(
REFERER,
HeaderValue::from_static("https://speed.cloudflare.com/"),
);
}
let client = reqwest::Client::builder()
.pool_max_idle_per_host(10)
.default_headers(headers)
.build()?;
let download_url = if url.contains("cloudflare") {
format!("{}?bytes={}", url, 100_000_000)
} else {
url.to_string()
};
let start = Instant::now();
let total_bytes = Arc::new(AtomicU64::new(0));
let duration = std::time::Duration::from_secs(duration_secs);
let num_streams = 4;
let mut handles = Vec::new();
for _ in 0..num_streams {
let client = client.clone();
let url = download_url.clone();
let bytes_counter = Arc::clone(&total_bytes);
let test_duration = duration;
let test_start = start;
let handle = tokio::spawn(async move {
while test_start.elapsed() < test_duration {
match client.get(&url).send().await {
Ok(response) if response.status().is_success() => {
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
if test_start.elapsed() >= test_duration {
break;
}
if let Ok(data) = chunk {
bytes_counter.fetch_add(data.len() as u64, Ordering::Relaxed);
}
}
}
_ => {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
}
}
});
handles.push(handle);
}
let mut samples: Vec<SpeedSample> = Vec::new();
let mut peak_speed: f64 = 0.0;
let mut last_bytes: u64 = 0;
let mut last_time = Instant::now();
while start.elapsed() < duration {
tokio::time::sleep(tokio::time::Duration::from_millis(SAMPLE_INTERVAL_MS)).await;
let current_bytes = total_bytes.load(Ordering::Relaxed);
let elapsed_since_last = last_time.elapsed().as_secs_f64();
let bytes_diff = current_bytes.saturating_sub(last_bytes);
let speed_mbps = if elapsed_since_last > 0.0 {
format_speed_mbps(bytes_diff as f64 / elapsed_since_last)
} else {
0.0
};
if speed_mbps > peak_speed {
peak_speed = speed_mbps;
}
samples.push(SpeedSample {
timestamp_ms: start.elapsed().as_millis() as u64,
speed_mbps,
});
let total_elapsed = start.elapsed().as_secs_f64();
let avg_speed = if total_elapsed > 0.0 {
format_speed_mbps(current_bytes as f64 / total_elapsed)
} else {
0.0
};
progress_callback(SpeedProgress {
current_speed_mbps: speed_mbps,
average_speed_mbps: avg_speed,
bytes_transferred: current_bytes,
elapsed_secs: total_elapsed,
progress_percent: (total_elapsed / duration_secs as f64) * 100.0,
});
last_bytes = current_bytes;
last_time = Instant::now();
}
for handle in handles {
handle.abort();
}
let final_bytes = total_bytes.load(Ordering::Relaxed);
let total_duration = start.elapsed().as_secs_f64();
let average_speed = if total_duration > 0.0 {
format_speed_mbps(final_bytes as f64 / total_duration)
} else {
0.0
};
Ok(SpeedResult {
speed_mbps: average_speed,
peak_speed_mbps: peak_speed,
bytes_transferred: final_bytes,
duration_secs: total_duration,
samples,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_download_cloudflare() {
let result = measure_download("https://speed.cloudflare.com/__down", 3, |_| {}).await;
assert!(result.is_ok());
}
}