From f9fc87c664578aade05db708774730dff005235f Mon Sep 17 00:00:00 2001 From: Lowder Date: Wed, 23 Sep 2026 10:24:39 +0500 Subject: [PATCH] feat: diagnostics commands --- Cargo.lock | 4 +- probe/Cargo.toml | 2 +- probe/src/main.rs | 141 +++++++++++++++++++++++++++++++----- probe/src/traceroute.rs | 151 ++++++++++++++++++++++++++++++++++++--- reports/src/probe.rs | 47 ++++++++++++ website/Cargo.toml | 2 +- website/src/mqtt_auth.rs | 14 ++++ 7 files changed, 331 insertions(+), 30 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 3efd2ff..3f9053d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2488,7 +2488,7 @@ dependencies = [ [[package]] name = "probe" -version = "0.6.5" +version = "0.6.6" dependencies = [ "anyhow", "clap", @@ -4509,7 +4509,7 @@ dependencies = [ [[package]] name = "website" -version = "1.3.6" +version = "1.3.7" dependencies = [ "anyhow", "dotenvy", diff --git a/probe/Cargo.toml b/probe/Cargo.toml index e18d944..495fbfb 100644 --- a/probe/Cargo.toml +++ b/probe/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "probe" -version = "0.6.5" +version = "0.6.6" edition = "2024" license-file = "../LICENSE" description = "Dynamic network probe daemon for Cheburcheck" diff --git a/probe/src/main.rs b/probe/src/main.rs index 6d80f1c..ecb7735 100644 --- a/probe/src/main.rs +++ b/probe/src/main.rs @@ -7,7 +7,10 @@ mod update; use anyhow::{Context, Result, bail}; use clap::{Parser, Subcommand}; use log::{debug, error, info, warn}; -use reports::probe::{DpiProbeConfig, ProbeConfig, ProbeResult, ProbeStatus, ProbeTask}; +use reports::probe::{ + DpiProbeConfig, ProbeCommand, ProbeCommandResult, ProbeConfig, ProbeResult, ProbeStatus, + ProbeTask, +}; use rumqttc::{ AsyncClient, Event, Incoming, LastWill, MqttOptions, NetworkOptions, QoS, Transport, }; @@ -19,6 +22,7 @@ use tokio::sync::RwLock; const CONFIG_TOPIC: &str = "probe/config/v1"; const UPDATE_TOPIC: &str = "probe/update/v1"; +const PUBLIC_TASK_TOPIC: &str = "probe/tasks/v1/+"; const UPDATE_REQUEST_COMMAND: &str = "/usr/libexec/cheburprobe-request-update"; const MQTT_MAX_PACKET_SIZE: usize = 1024 * 1024; @@ -162,14 +166,9 @@ async fn main() -> Result<()> { } else { debug!("MQTT-triggered updates are disabled for this standalone installation"); } + subscribe_to_tasks(&client, &args.probe_id).await?; client - .subscribe("probe/tasks/v1/+", QoS::AtLeastOnce) - .await?; - client - .subscribe( - format!("probe/tasks/v1/{}/+", args.probe_id), - QoS::AtLeastOnce, - ) + .subscribe(command_subscription(&args.probe_id), QoS::AtLeastOnce) .await?; info!( @@ -194,7 +193,65 @@ async fn main() -> Result<()> { config.clone(), publish.payload.to_vec(), ); - } else { + } else if let Some(command_id) = command_id(&publish.topic, &args.probe_id) { + if publish.retain { + warn!("ignoring retained command on {}", publish.topic); + continue; + } + let command_id = command_id.to_owned(); + let client = client.clone(); + let probe_id = args.probe_id.clone(); + let payload = publish.payload.to_vec(); + let retries = args.traceroute_retries; + let semaphore = task_semaphore.clone(); + tokio::spawn(async move { + let result = match serde_json::from_slice::(&payload) { + Ok(ProbeCommand::ResubscribeTasks) => { + match resubscribe_to_tasks(&client).await { + Ok(()) => { + ProbeCommandResult::ResubscribeTasks { requested: true } + } + Err(error) => ProbeCommandResult::Error { + message: error.to_string(), + }, + } + } + Ok(ProbeCommand::Traceroute { target, max_hops }) => { + match semaphore.acquire_owned().await { + Ok(_permit) => match traceroute::manual_traceroute( + target, max_hops, retries, + ) + .await + { + Ok(hops) => ProbeCommandResult::Traceroute { target, hops }, + Err(error) => ProbeCommandResult::Error { + message: error.to_string(), + }, + }, + Err(error) => ProbeCommandResult::Error { + message: error.to_string(), + }, + } + } + Err(error) => ProbeCommandResult::Error { + message: format!("invalid command: {error}"), + }, + }; + let result_topic = + format!("probe/command-results/v1/{probe_id}/{command_id}"); + if let Err(error) = client + .publish( + result_topic, + QoS::AtLeastOnce, + false, + serde_json::to_vec(&result).expect("serialize command result"), + ) + .await + { + warn!("failed to publish command result: {error}"); + } + }); + } else if probe_task_job_id(&publish.topic).is_some() { let client = client.clone(); let args = args.clone(); let config = config.clone(); @@ -249,14 +306,9 @@ async fn main() -> Result<()> { if mqtt_updates_enabled { subscribe_to_update_requests(&client, &args.probe_id).await?; } + subscribe_to_tasks(&client, &args.probe_id).await?; client - .subscribe("probe/tasks/v1/+", QoS::AtLeastOnce) - .await?; - client - .subscribe( - format!("probe/tasks/v1/{}/+", args.probe_id), - QoS::AtLeastOnce, - ) + .subscribe(command_subscription(&args.probe_id), QoS::AtLeastOnce) .await?; tokio::time::sleep(Duration::from_secs(2)).await; } @@ -264,6 +316,39 @@ async fn main() -> Result<()> { } } +fn command_subscription(probe_id: &str) -> String { + format!("probe/commands/v1/{probe_id}/+") +} + +fn command_id<'a>(topic: &'a str, probe_id: &str) -> Option<&'a str> { + match topic.split('/').collect::>().as_slice() { + ["probe", "commands", "v1", recipient, command_id] + if *recipient == probe_id && !command_id.is_empty() => + { + Some(command_id) + } + _ => None, + } +} + +async fn subscribe_to_tasks(client: &AsyncClient, probe_id: &str) -> Result<()> { + client + .subscribe(PUBLIC_TASK_TOPIC, QoS::AtLeastOnce) + .await?; + client + .subscribe(format!("probe/tasks/v1/{probe_id}/+"), QoS::AtLeastOnce) + .await?; + Ok(()) +} + +async fn resubscribe_to_tasks(client: &AsyncClient) -> Result<()> { + client.unsubscribe(PUBLIC_TASK_TOPIC).await?; + client + .subscribe(PUBLIC_TASK_TOPIC, QoS::AtLeastOnce) + .await?; + Ok(()) +} + fn spawn_config_update( client: AsyncClient, status_topic: String, @@ -599,6 +684,30 @@ mod tests { assert!(!is_update_topic("probe/update/v1/42/extra", "42")); } + #[test] + fn command_topics_are_node_scoped() { + assert_eq!(command_subscription("42"), "probe/commands/v1/42/+"); + assert_eq!( + command_id("probe/commands/v1/42/trace-1", "42"), + Some("trace-1") + ); + assert_eq!(command_id("probe/commands/v1/7/trace-1", "42"), None); + assert_eq!(command_id("probe/commands/v1/42/trace-1/extra", "42"), None); + } + + #[test] + fn manual_commands_decode() { + let command: ProbeCommand = + serde_json::from_str(r#"{"type":"traceroute","target":"1.1.1.1"}"#).unwrap(); + assert!(matches!( + command, + ProbeCommand::Traceroute { max_hops: 30, .. } + )); + let command: ProbeCommand = + serde_json::from_str(r#"{"type":"resubscribe_tasks"}"#).unwrap(); + assert!(matches!(command, ProbeCommand::ResubscribeTasks)); + } + #[test] fn decodes_separate_dpi_targets() { let config: ProbeConfig = serde_json::from_value(serde_json::json!({ diff --git a/probe/src/traceroute.rs b/probe/src/traceroute.rs index 3647d6c..7c6f359 100644 --- a/probe/src/traceroute.rs +++ b/probe/src/traceroute.rs @@ -2,16 +2,141 @@ use etherparse::{ Icmpv4Type, Icmpv6Slice, Icmpv6Type, IpNumber, LaxNetSlice, LaxSlicedPacket, TransportSlice, icmpv4, icmpv6, }; +use hickory_resolver::proto::rr::RData; use polling::{Event, Events, Poller}; -use reports::probe::{TcpTracerouteOutcome, TcpTracerouteResult}; +use reports::probe::{ + ManualTracerouteHop, ManualTracerouteOutcome, TcpTracerouteOutcome, TcpTracerouteResult, +}; use socket2::{Domain, Protocol, SockAddr, Socket, Type}; -use std::io::{self, Read}; +use std::io; +use std::mem::MaybeUninit; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::time::{Duration, Instant}; const HTTPS_PORT: u16 = 443; const HOP_TIMEOUT: Duration = Duration::from_secs(1); +pub async fn manual_traceroute( + target: IpAddr, + max_hops: u8, + retries: u8, +) -> io::Result> { + let mut hops = + tokio::task::spawn_blocking(move || trace_all_blocking(target, max_hops, retries)) + .await + .map_err(io::Error::other)??; + let resolver = + match hickory_resolver::Resolver::builder_tokio().and_then(|builder| builder.build()) { + Ok(resolver) => resolver, + Err(error) => { + log::warn!("reverse DNS resolver unavailable: {error}"); + return Ok(hops); + } + }; + let names = futures::future::join_all(hops.iter().map(|hop| async { + let Some(address) = hop.address else { + return Vec::new(); + }; + match tokio::time::timeout(Duration::from_secs(3), resolver.reverse_lookup(address)).await { + Ok(Ok(names)) => names + .answers() + .iter() + .filter_map(|record| match &record.data { + RData::PTR(name) => Some(name.to_string().trim_end_matches('.').to_string()), + _ => None, + }) + .collect(), + _ => Vec::new(), + } + })) + .await; + for (hop, names) in hops.iter_mut().zip(names) { + hop.reverse_names = names; + } + Ok(hops) +} + +fn trace_all_blocking( + target: IpAddr, + max_hops: u8, + retries: u8, +) -> io::Result> { + if max_hops == 0 || max_hops > 64 || retries == 0 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "max_hops must be 1..=64 and retries must be nonzero", + )); + } + let (domain, protocol) = match target { + IpAddr::V4(_) => (Domain::IPV4, Protocol::ICMPV4), + IpAddr::V6(_) => (Domain::IPV6, Protocol::ICMPV6), + }; + let receiver = Socket::new(domain, Type::RAW, Some(protocol))?; + let mut hops = Vec::new(); + for ttl in 1..=max_hops { + let attempts = connect_attempts(target, ttl, retries)?; + let response = if attempts.is_empty() { + HopResponse::Timeout + } else { + wait_for_hop_response(&receiver, &attempts, target)? + }; + let (address, outcome, complete) = match response { + HopResponse::IcmpTimeExceeded(address) => ( + Some(address), + ManualTracerouteOutcome::IcmpTimeExceeded, + false, + ), + HopResponse::Rst => (Some(target), ManualTracerouteOutcome::Rst, true), + HopResponse::Connected => (Some(target), ManualTracerouteOutcome::Connected, true), + HopResponse::Timeout => (None, ManualTracerouteOutcome::Timeout, false), + }; + hops.push(ManualTracerouteHop { + ttl, + address, + reverse_names: Vec::new(), + outcome, + }); + if complete { + break; + } + } + Ok(hops) +} + +fn connect_attempts(target: IpAddr, ttl: u8, retries: u8) -> io::Result> { + let domain = match target { + IpAddr::V4(_) => Domain::IPV4, + IpAddr::V6(_) => Domain::IPV6, + }; + let destination = SockAddr::from(SocketAddr::new(target, HTTPS_PORT)); + let mut attempts = Vec::with_capacity(retries as usize); + for _ in 0..retries { + let tcp = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?; + tcp.set_nonblocking(true)?; + match target { + IpAddr::V4(_) => tcp.set_ttl_v4(ttl as u32)?, + IpAddr::V6(_) => tcp.set_unicast_hops_v6(ttl as u32)?, + } + let unspecified = match target { + IpAddr::V4(_) => IpAddr::V4(Ipv4Addr::UNSPECIFIED), + IpAddr::V6(_) => IpAddr::V6(Ipv6Addr::UNSPECIFIED), + }; + tcp.bind(&SockAddr::from(SocketAddr::new(unspecified, 0)))?; + if let Err(error) = tcp.connect(&destination) { + if !is_connect_in_progress(&error) && error.kind() != io::ErrorKind::ConnectionRefused { + continue; + } + } + let source_port = tcp + .local_addr()? + .as_socket() + .map(|addr| addr.port()) + .unwrap_or(0); + attempts.push((tcp, source_port)); + } + Ok(attempts) +} + pub async fn tcp_traceroute( target: IpAddr, start_hop: u8, @@ -82,7 +207,7 @@ fn trace_blocking( return Err(last_error.unwrap_or_else(|| io::Error::other("no traceroute attempts"))); } match wait_for_hop_response(&receiver, &tcp_attempts, target)? { - HopResponse::IcmpTimeExceeded => { + HopResponse::IcmpTimeExceeded(_) => { return Ok(TcpTracerouteResult { target, result: TcpTracerouteOutcome::IcmpTimeExceeded { hop: ttl }, @@ -120,7 +245,7 @@ fn is_connect_in_progress(error: &io::Error) -> bool { } enum HopResponse { - IcmpTimeExceeded, + IcmpTimeExceeded(IpAddr), Rst, Connected, Timeout, @@ -166,7 +291,7 @@ fn wait_on_poller( let deadline = Instant::now() + HOP_TIMEOUT; let mut watching_tcp = vec![true; tcp_attempts.len()]; - let mut buffer = [0u8; 2048]; + let mut buffer = [MaybeUninit::::uninit(); 2048]; let mut events = Events::new(); loop { @@ -197,14 +322,20 @@ fn wait_on_poller( } if events.iter().any(|event| event.key == ICMP_KEY) { - let mut raw = receiver; - match raw.read(&mut buffer) { - Ok(bytes) + match receiver.recv_from(&mut buffer) { + Ok((bytes, source)) if tcp_attempts.iter().any(|(_, source_port)| { - is_matching_time_exceeded(&buffer[..bytes], target, *source_port) + // SAFETY: recv_from initialized the first `bytes` elements. + let packet = unsafe { + std::slice::from_raw_parts(buffer.as_ptr().cast::(), bytes) + }; + is_matching_time_exceeded(packet, target, *source_port) }) => { - return Ok(HopResponse::IcmpTimeExceeded); + if let Some(address) = source.as_socket().map(|addr| addr.ip()) { + return Ok(HopResponse::IcmpTimeExceeded(address)); + } + poller.modify(receiver, Event::readable(ICMP_KEY))?; } Ok(_) => poller.modify(receiver, Event::readable(ICMP_KEY))?, Err(error) if error.kind() == io::ErrorKind::WouldBlock => { diff --git a/reports/src/probe.rs b/reports/src/probe.rs index 982801a..e90c2f8 100644 --- a/reports/src/probe.rs +++ b/reports/src/probe.rs @@ -152,6 +152,53 @@ pub struct TcpTracerouteResult { pub result: TcpTracerouteOutcome, } +#[derive(Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ProbeCommand { + Traceroute { + target: IpAddr, + #[serde(default = "default_manual_max_hops")] + max_hops: u8, + }, + ResubscribeTasks, +} + +pub const fn default_manual_max_hops() -> u8 { + 30 +} + +#[derive(Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ProbeCommandResult { + Traceroute { + target: IpAddr, + hops: Vec, + }, + ResubscribeTasks { + requested: bool, + }, + Error { + message: String, + }, +} + +#[derive(Clone, Serialize, Deserialize)] +pub struct ManualTracerouteHop { + pub ttl: u8, + pub address: Option, + pub reverse_names: Vec, + pub outcome: ManualTracerouteOutcome, +} + +#[derive(Clone, Copy, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ManualTracerouteOutcome { + IcmpTimeExceeded, + Rst, + Connected, + Timeout, +} + #[derive(Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum TcpTracerouteOutcome { diff --git a/website/Cargo.toml b/website/Cargo.toml index 60b722e..ec7f0b6 100644 --- a/website/Cargo.toml +++ b/website/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "website" -version = "1.3.6" +version = "1.3.7" edition = "2024" [dependencies] diff --git a/website/src/mqtt_auth.rs b/website/src/mqtt_auth.rs index 1bb3806..92f9ebf 100644 --- a/website/src/mqtt_auth.rs +++ b/website/src/mqtt_auth.rs @@ -123,6 +123,7 @@ async fn can_probe_subscribe(client_id: &str, topic: &str, pool: &PgPool) -> boo if topic == "probe/config/v1" || topic == "probe/update/v1" || topic == format!("probe/update/v1/{client_id}") + || topic == format!("probe/commands/v1/{client_id}/+") || is_own_task_subscription(client_id, topic) { return true; @@ -157,6 +158,11 @@ fn can_probe_publish(client_id: &str, topic: &str) -> bool { return true; } + if let Some(command_id) = topic.strip_prefix(&format!("probe/command-results/v1/{client_id}/")) + { + return !command_id.is_empty() && !command_id.contains('/'); + } + let mut parts = topic.split('/'); matches!( ( @@ -192,11 +198,19 @@ mod tests { assert!(can_probe_subscribe("42", "probe/update/v1", &pool).await); assert!(can_probe_subscribe("42", "probe/update/v1/42", &pool).await); assert!(!can_probe_subscribe("42", "probe/update/v1/7", &pool).await); + assert!(can_probe_subscribe("42", "probe/commands/v1/42/+", &pool).await); + assert!(!can_probe_subscribe("42", "probe/commands/v1/7/+", &pool).await); } #[test] fn results_can_only_be_published_as_the_authenticated_node() { assert!(can_probe_publish("42", "probe/results/v1/job/42")); assert!(!can_probe_publish("42", "probe/results/v1/job/7")); + assert!(can_probe_publish("42", "probe/command-results/v1/42/job")); + assert!(!can_probe_publish("42", "probe/command-results/v1/7/job")); + assert!(!can_probe_publish( + "42", + "probe/command-results/v1/42/job/extra" + )); } }