1use std::sync::Arc;
89
90use cf_mach::{
91 nq_core::{ConnectionType, Network, Time, TokioTime},
92 nq_latency::{Latency, LatencyConfig, LatencyResult},
93 nq_rpm::{ConnectionErrorPolicy, Responsiveness, ResponsivenessConfig, ResponsivenessResult},
94 nq_tokio_network::TokioNetwork,
95};
96use reqwest::Url;
97use serde::{Deserialize, Deserializer};
98use tokio_util::sync::CancellationToken;
99
100use super::prelude::*;
101
102make_log_macro!(debug, "speedtest");
103
104#[derive(Deserialize, Debug, SmartDefault)]
105#[serde(deny_unknown_fields, default)]
106pub struct Config {
107 pub format: FormatConfig,
108 #[default(1800.into())]
109 pub interval: Seconds,
110 #[serde(flatten)]
111 pub speedtest: SpeedtestConfig,
112}
113
114#[derive(Debug, Deserialize, SmartDefault)]
115#[serde(deny_unknown_fields, default)]
116pub struct SpeedtestConfig {
117 #[serde(deserialize_with = "deserialize_url_opt")]
132 pub config_url: Option<Url>,
133 #[default("https://h3.speed.cloudflare.com/__down?bytes=10000000000".parse().unwrap())]
135 pub large_download_url: Url,
136 #[default("https://h3.speed.cloudflare.com/__down?bytes=10".parse().unwrap())]
138 #[serde(deserialize_with = "deserialize_url")]
139 pub small_download_url: Url,
140 #[default("https://h3.speed.cloudflare.com/__up".parse().unwrap())]
142 #[serde(deserialize_with = "deserialize_url")]
143 pub upload_url: Url,
144 pub latency: LatencyConfigOpts,
145 pub rpm: RpmConfigOpts,
146}
147
148#[derive(Debug, Deserialize, SmartDefault)]
149#[serde(deny_unknown_fields, default)]
150pub struct LatencyConfigOpts {
151 #[default(20)]
153 pub runs: usize,
154}
155
156#[derive(Debug, Deserialize, SmartDefault)]
157#[serde(deny_unknown_fields, default)]
158pub struct RpmConfigOpts {
159 #[default(4)]
161 pub moving_average_distance: usize,
162 #[default(0.05)]
165 pub std_tolerance: f64,
166 #[default(0.95)]
170 pub trimmed_mean_percent: f64,
171 #[default(16)]
174 pub max_loaded_connections: usize,
175 #[default(Duration::from_millis(500))]
177 #[serde(deserialize_with = "deserialize_duration_ms")]
178 pub interval_duration_ms: Duration,
179 #[default(Duration::from_millis(12_000))]
181 #[serde(deserialize_with = "deserialize_duration_ms")]
182 pub test_duration_ms: Duration,
183 #[default(ConnectionType::H2)]
185 #[serde(deserialize_with = "deserialize_conn_type")]
186 pub conn_type: ConnectionType,
187 #[default(100_000_000)]
199 pub upload_bytes_per_request: usize,
200}
201
202pub(crate) fn prepare(config: &Config) -> Result<Arc<BlockPlan>> {
203 BlockPlan::new(vec![OutputPlan::new(
206 "main",
207 config.format.with_default(
208 " ^icon_ping $ping.eng(prefix:m) ^icon_net_down $speed_down ^icon_net_up $speed_up ",
209 )?,
210 )])
211}
212
213pub(crate) async fn run(config: &Config, api: &CommonApi, plan: &Arc<BlockPlan>) -> Result<()> {
214 let output_main = plan.output("main")?;
215 let format = output_main.format();
216
217 let need_ping = format.contains_key("ping");
218 let need_jitter = format.contains_key("jitter");
219 let need_speed_down = format.contains_key("speed_down");
220 let need_speed_up = format.contains_key("speed_up");
221
222 loop {
223 let speedtest_urls = get_speedtest_urls(&config.speedtest).await?;
224
225 let mut values = HashMap::new();
226
227 if need_ping || need_jitter {
228 debug!("running latency test");
229
230 let latency_results = test_latency(LatencyConfig {
231 url: speedtest_urls.small_https_download_url.clone(),
232 runs: config.speedtest.latency.runs,
233 scoped_headers: None,
234 })
235 .await?;
236
237 if need_ping {
238 values.insert(
239 "ping".into(),
240 Value::seconds(latency_results.median().error("no median RTT available")?),
241 );
242 }
243
244 if need_jitter {
245 values.insert(
246 "jitter".into(),
247 Value::seconds(latency_results.jitter().error("no jitter available")?),
248 );
249 }
250 }
251
252 if need_speed_down || need_speed_up {
253 let responsiveness_config = ResponsivenessConfig {
254 large_download_url: speedtest_urls.large_https_download_url,
255 small_download_url: speedtest_urls.small_https_download_url,
256 upload_url: speedtest_urls.https_upload_url,
257 moving_average_distance: config.speedtest.rpm.moving_average_distance,
258 interval_duration: config.speedtest.rpm.interval_duration_ms,
259 test_duration: config.speedtest.rpm.test_duration_ms,
260 trimmed_mean_percent: config.speedtest.rpm.trimmed_mean_percent,
261 std_tolerance: config.speedtest.rpm.std_tolerance,
262 max_loaded_connections: config.speedtest.rpm.max_loaded_connections,
263 conn_type: config.speedtest.rpm.conn_type,
264 upload_bytes_per_request: config.speedtest.rpm.upload_bytes_per_request,
265 determine_load_only: false,
267 on_connection_error: ConnectionErrorPolicy::default(),
268 scoped_headers: None,
269 };
270
271 if need_speed_down {
272 debug!("running download test");
273 let download_result = test_network_speed(&responsiveness_config, true).await?;
274 values.insert(
275 "speed_down".into(),
276 Value::bits(
277 download_result
278 .throughput()
279 .error("no download throughput available")?,
280 ),
281 );
282 }
283
284 if need_speed_up {
285 debug!("running upload test");
286 let upload_result = test_network_speed(&responsiveness_config, false).await?;
287 values.insert(
288 "speed_up".into(),
289 Value::bits(
290 upload_result
291 .throughput()
292 .error("no upload throughput available")?,
293 ),
294 );
295 }
296 }
297
298 let mut widget = output_main.new_widget();
299 widget.set_values(values);
300 api.set_widget(widget)?;
301
302 select! {
303 _ = sleep(config.interval.0) => (),
304 _ = api.wait_for_update_request() => (),
305 }
306 }
307}
308
309#[derive(Debug, Deserialize)]
310struct SpeedtestUrls {
311 #[serde(alias = "small_download_url", deserialize_with = "deserialize_url")]
312 small_https_download_url: Url,
313 #[serde(alias = "large_download_url", deserialize_with = "deserialize_url")]
314 large_https_download_url: Url,
315 #[serde(alias = "upload_url", deserialize_with = "deserialize_url")]
316 https_upload_url: Url,
317}
318
319#[derive(Deserialize)]
320struct RpmServerConfig {
321 urls: SpeedtestUrls,
322}
323
324async fn get_speedtest_urls(speedtest_config: &SpeedtestConfig) -> Result<SpeedtestUrls> {
326 match speedtest_config.config_url.clone() {
327 Some(config_url) => {
328 debug!("fetching configuration from {config_url}");
329 let urls = REQWEST_CLIENT
330 .get(config_url)
331 .send()
332 .await
333 .error("Failed to send request with reqwest")?
334 .json::<RpmServerConfig>()
335 .await
336 .error("Failed to parse JSON from rpm config endpoint")?
337 .urls;
338 debug!("retrieved configuration urls: {urls:?}");
339
340 Ok(urls)
341 }
342 None => Ok(SpeedtestUrls {
343 small_https_download_url: speedtest_config.small_download_url.clone(),
344 large_https_download_url: speedtest_config.large_download_url.clone(),
345 https_upload_url: speedtest_config.upload_url.clone(),
346 }),
347 }
348}
349
350async fn test_latency(config: LatencyConfig) -> Result<LatencyResult> {
351 let shutdown = CancellationToken::new();
352 let time: Arc<dyn Time> = Arc::new(TokioTime::new());
353 let network: Arc<dyn Network> =
354 Arc::new(TokioNetwork::new(Arc::clone(&time), shutdown.clone()));
355
356 let rtt = Latency::new(config);
357 let result = rtt
358 .run_test(network, time, shutdown.clone())
359 .await
360 .map_err(|e| Error::new(e.to_string()))?;
361
362 debug!("shutting down latency test");
363 let _ = tokio::time::timeout(tokio::time::Duration::from_secs(1), async {
364 shutdown.cancel();
365 })
366 .await;
367
368 Ok(result)
369}
370
371async fn test_network_speed(
372 config: &ResponsivenessConfig,
373 download: bool,
374) -> Result<ResponsivenessResult> {
375 let shutdown = CancellationToken::new();
376 let time: Arc<dyn Time> = Arc::new(TokioTime::new());
377 let network: Arc<dyn Network> =
378 Arc::new(TokioNetwork::new(Arc::clone(&time), shutdown.clone()));
379
380 let rpm =
381 Responsiveness::new(config.clone(), download).map_err(|e| Error::new(e.to_string()))?;
382 let result = rpm
383 .run_test(network, time, shutdown.clone())
384 .await
385 .map_err(|e| Error::new(e.to_string()))?;
386
387 debug!("shutting down network speed test");
388 let _ = tokio::time::timeout(tokio::time::Duration::from_secs(1), async {
389 shutdown.cancel();
390 })
391 .await;
392
393 Ok(result)
394}
395
396fn deserialize_url_opt<'de, D>(deserializer: D) -> Result<Option<Url>, D::Error>
397where
398 D: Deserializer<'de>,
399{
400 let url_opt = Option::<String>::deserialize(deserializer)?;
401 url_opt
402 .map(|url| url.parse().map_err(serde::de::Error::custom))
403 .transpose()
404}
405
406fn deserialize_url<'de, D>(deserializer: D) -> Result<Url, D::Error>
407where
408 D: Deserializer<'de>,
409{
410 let url = String::deserialize(deserializer)?;
411 url.parse().map_err(serde::de::Error::custom)
412}
413
414fn deserialize_duration_ms<'de, D>(deserializer: D) -> Result<Duration, D::Error>
415where
416 D: Deserializer<'de>,
417{
418 let duration_ms = u64::deserialize(deserializer)?;
419 Ok(Duration::from_millis(duration_ms))
420}
421
422fn deserialize_conn_type<'de, D>(deserializer: D) -> Result<ConnectionType, D::Error>
423where
424 D: Deserializer<'de>,
425{
426 let conn_type_str = String::deserialize(deserializer)?;
427 match conn_type_str.as_str() {
428 "h1" => Ok(ConnectionType::H1 { use_tls: true }),
429 "h2" => Ok(ConnectionType::H2),
430 "h3" => Ok(ConnectionType::H3),
431 _ => Err(serde::de::Error::custom(format!(
432 "Invalid connection type: {}. Must be one of: h1, h2, h3",
433 conn_type_str
434 ))),
435 }
436}
437
438#[cfg(test)]
439mod tests {
440 use super::*;
441
442 #[test]
443 fn plan_declares_main_output_without_icon_placeholders() {
444 let plan = prepare(&Config::default()).unwrap();
445 let declared: Vec<_> = plan.outputs().map(|o| o.id()).collect();
446 assert_eq!(declared, ["main"]);
447 let output = plan.output("main").unwrap();
448 assert_eq!(output.output().icon_placeholders().count(), 0);
451 assert_eq!(
452 output.output().static_icons(),
453 ["ping", "net_down", "net_up"]
454 );
455 }
456
457 #[test]
458 fn a_format_without_icons_declares_none() {
459 let config = Config {
460 format: " $ping ".parse().unwrap(),
461 ..Config::default()
462 };
463 let plan = prepare(&config).unwrap();
464 let output = plan.output("main").unwrap();
465 assert!(output.output().static_icons().is_empty());
466 }
467
468 #[test]
469 fn custom_format_is_respected() {
470 let config = Config {
471 format: " $ping ".parse().unwrap(),
472 ..Config::default()
473 };
474 let plan = prepare(&config).unwrap();
475 let output = plan.output("main").unwrap();
476 assert!(output.format().contains_key("ping"));
477 assert!(!output.format().contains_key("speed_down"));
478 }
479}