cf-mach 0.2.2

Network quality measurement CLI for latency, throughput, packet loss, and responsiveness
// Copyright (c) 2023-2024 Cloudflare, Inc.
// Licensed under the BSD-3-Clause license found in the LICENSE file or at https://opensource.org/licenses/BSD-3-Clause

use std::{fmt::Debug, sync::Arc};

use crate::nq_core::{
    ConnectionType, Network, ScopedHeaders, Time, Timestamp, client::MACH_USER_AGENT,
};
use crate::nq_stats::TimeSeries;
use anyhow::Context;
use http::{HeaderValue, Request, header::USER_AGENT};
use http_body_util::BodyExt;
use tokio_util::sync::CancellationToken;
use tracing::info;
use url::Url;

#[derive(Debug, Clone)]
pub struct LatencyConfig {
    pub url: Url,
    pub runs: usize,
    /// Headers attached only to requests whose host matches the scope's
    /// allowlist.
    pub scoped_headers: Option<ScopedHeaders>,
}

impl Default for LatencyConfig {
    fn default() -> Self {
        Self {
            url: "https://h3.speed.cloudflare.com/__down?bytes=10"
                .parse()
                .unwrap(),
            runs: 20,
            scoped_headers: None,
        }
    }
}

pub struct Latency {
    start: Timestamp,
    config: LatencyConfig,
    probe_results: TimeSeries,
}

impl Latency {
    pub fn new(config: LatencyConfig) -> Self {
        Self {
            start: Timestamp::now(),
            config,
            probe_results: TimeSeries::new(),
        }
    }

    pub async fn run_test(
        mut self,
        network: Arc<dyn Network>,
        time: Arc<dyn Time>,
        _shutdown: CancellationToken,
    ) -> anyhow::Result<LatencyResult> {
        self.start = time.now();

        for run in 0..self.config.runs {
            let url = self.config.url.to_owned();
            let network = Arc::clone(&network);
            let time = Arc::clone(&time);

            let host = url
                .host_str()
                .context("small download url must have a domain")?;
            let host_with_port = format!("{}:{}", host, url.port_or_known_default().unwrap_or(443));

            let conn_start = time.now();

            let addrs = network
                .resolve(host_with_port)
                .await
                .context("unable to resolve host")?;
            let time_lookup = time.now();

            let conn_type = ConnectionType::H1 { use_tls: true };
            let connection = network
                .new_connection(conn_start, addrs[0], host.to_string(), conn_type)
                .await
                .context("unable to create new connection")?;
            {
                let conn = connection.write().await;
                conn.timing().set_lookup(time_lookup);
            }

            let tcp_handshake_duration = {
                let conn = connection.read().await;
                conn.timing()
                    .time_connect()
                    .saturating_sub(conn.timing().time_lookup())
            };

            info!(
                "latency run {run}: {:2.4} s.",
                tcp_handshake_duration.as_secs_f32()
            );

            // perform a simple GET to do some amount of work
            let mut request = Request::get(url.as_str()).body(Default::default())?;
            request
                .headers_mut()
                .insert(USER_AGENT, HeaderValue::from_static(MACH_USER_AGENT));
            if let Some(scoped_headers) = &self.config.scoped_headers {
                let request_uri = request.uri().clone();
                scoped_headers.apply(&request_uri, request.headers_mut());
            }

            let response = network
                .send_request(connection, request)
                .await
                .context("GET request failed")?;

            let _ = response
                .into_body()
                .collect()
                .await
                .context("unable to read GET request body")?;

            self.probe_results
                .add(conn_start, tcp_handshake_duration.as_secs_f64());
        }

        // while let Some(res) = task_set.join_next().await {
        //     let (conn_start, tcp_handshake_duration) = res??;
        //     self.probe_results.add(conn_start, tcp_handshake_duration);
        // }

        Ok(LatencyResult {
            measurements: self.probe_results,
        })
    }
}

#[derive(Default, Debug)]
pub struct LatencyResult {
    pub measurements: TimeSeries,
}

impl LatencyResult {
    pub fn median(&self) -> Option<f64> {
        self.measurements.quantile(0.50)
    }

    /// Jitter as calculated as the average distance between consecutive rtt
    /// measurments.
    pub fn jitter(&self) -> Option<f64> {
        let values: Vec<_> = self.measurements.values().collect();

        let distances = values.windows(2).filter_map(|window| {
            window
                .last()
                .zip(window.first())
                .map(|(last, first)| (last - first).abs())
        });

        let sum: f64 = distances.clone().sum();
        let count = distances.count();

        if count > 0 {
            Some(sum / count as f64)
        } else {
            None
        }
    }
}