speedtest-tui 0.1.1

A terminal-based network speed test tool with real-time gauges and graphs
use anyhow::Result;
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 CHUNK_SIZE: usize = 1_000_000;
const SAMPLE_INTERVAL_MS: u64 = 250;

pub async fn measure_upload<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 start = Instant::now();
    let total_bytes = Arc::new(AtomicU64::new(0));
    let duration = std::time::Duration::from_secs(duration_secs);
    let upload_data: Vec<u8> = vec![0u8; CHUNK_SIZE];
    let num_streams = 4;
    let mut handles = Vec::new();

    for _ in 0..num_streams {
        let client = client.clone();
        let url = url.to_string();
        let bytes_counter = Arc::clone(&total_bytes);
        let data = upload_data.clone();
        let test_duration = duration;
        let test_start = start;

        let handle = tokio::spawn(async move {
            while test_start.elapsed() < test_duration {
                match client
                    .post(&url)
                    .body(data.clone())
                    .header("Content-Type", "application/octet-stream")
                    .send()
                    .await
                {
                    Ok(response) if response.status().is_success() => {
                        bytes_counter.fetch_add(CHUNK_SIZE as u64, Ordering::Relaxed);
                    }
                    _ => {
                        tokio::time::sleep(tokio::time::Duration::from_millis(50)).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_upload_cloudflare() {
        let result = measure_upload("https://speed.cloudflare.com/__up", 3, |_| {}).await;
        assert!(result.is_ok());
    }
}