mirror of
https://github.com/LowderPlay/cheburcheck.git
synced 2026-10-11 16:48:15 +03:00
feat: ip probe traceroute
This commit is contained in:
+85
-176
@@ -1,24 +1,29 @@
|
||||
mod sni;
|
||||
mod traceroute;
|
||||
|
||||
use anyhow::{Context, Result, bail};
|
||||
use clap::Parser;
|
||||
use futures::future::join_all;
|
||||
use log::{error, info, warn};
|
||||
use reports::probe::{Host, HostProbeResult, ProbeConfig, ProbeEvidence, ProbeStatus, ProbeTask};
|
||||
use rand::seq::SliceRandom;
|
||||
use reports::probe::{ProbeConfig, ProbeResult, ProbeStatus, ProbeTask};
|
||||
use rumqttc::{
|
||||
AsyncClient, Event, Incoming, LastWill, MqttOptions, NetworkOptions, QoS, Transport,
|
||||
};
|
||||
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
|
||||
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
|
||||
use rustls::{ClientConfig, DigitallySignedStruct, Error as TlsError, SignatureScheme};
|
||||
use std::collections::HashSet;
|
||||
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio::sync::RwLock;
|
||||
use tokio::time;
|
||||
use tokio_rustls::TlsConnector;
|
||||
|
||||
const CONFIG_TOPIC: &str = "probe/config/v1";
|
||||
|
||||
#[derive(Clone)]
|
||||
struct LoadedProbeConfig {
|
||||
config: ProbeConfig,
|
||||
control_hosts_v4: Vec<Ipv4Addr>,
|
||||
control_hosts_v6: Vec<Ipv6Addr>,
|
||||
}
|
||||
|
||||
#[derive(Parser, Debug, Clone)]
|
||||
#[command(author, version, about = "Dynamic probing daemon")]
|
||||
struct Args {
|
||||
@@ -39,6 +44,9 @@ struct Args {
|
||||
|
||||
#[arg(long, env = "MAX_CONCURRENT_TASKS", default_value_t = 8)]
|
||||
max_concurrent_tasks: usize,
|
||||
|
||||
#[arg(long, env = "TRACEROUTE_MAX_HOPS", default_value_t = 5)]
|
||||
traceroute_max_hops: u8,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
@@ -48,6 +56,9 @@ async fn main() -> Result<()> {
|
||||
if args.max_concurrent_tasks == 0 {
|
||||
bail!("max_concurrent_tasks must be greater than zero");
|
||||
}
|
||||
if args.traceroute_max_hops == 0 {
|
||||
bail!("traceroute_max_hops must be greater than zero");
|
||||
}
|
||||
|
||||
let status_topic = format!("probe/status/v1/{}", args.probe_id);
|
||||
let offline_status = serde_json::to_vec(&ProbeStatus {
|
||||
@@ -155,9 +166,35 @@ fn mqtt_transport(mqtt_host: &str) -> Result<Transport> {
|
||||
}
|
||||
}
|
||||
|
||||
async fn update_config(config: &Arc<RwLock<Option<ProbeConfig>>>, payload: &[u8]) -> Result<()> {
|
||||
let value = serde_json::from_slice(payload).context("decode probe config")?;
|
||||
*config.write().await = Some(value);
|
||||
async fn update_config(
|
||||
config: &Arc<RwLock<Option<LoadedProbeConfig>>>,
|
||||
payload: &[u8],
|
||||
) -> Result<()> {
|
||||
let value: ProbeConfig = serde_json::from_slice(payload).context("decode probe config")?;
|
||||
let mut control_hosts_v4 = HashSet::new();
|
||||
let mut control_hosts_v6 = HashSet::new();
|
||||
for domain in &value.control_hosts {
|
||||
match tokio::net::lookup_host((domain.as_str(), 443)).await {
|
||||
Ok(addresses) => {
|
||||
for address in addresses {
|
||||
match address.ip() {
|
||||
IpAddr::V4(address) => {
|
||||
control_hosts_v4.insert(address);
|
||||
}
|
||||
IpAddr::V6(address) => {
|
||||
control_hosts_v6.insert(address);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => warn!("failed to resolve control host {domain}: {error}"),
|
||||
}
|
||||
}
|
||||
*config.write().await = Some(LoadedProbeConfig {
|
||||
config: value,
|
||||
control_hosts_v4: control_hosts_v4.into_iter().collect(),
|
||||
control_hosts_v6: control_hosts_v6.into_iter().collect(),
|
||||
});
|
||||
info!("updated retained probe config");
|
||||
Ok(())
|
||||
}
|
||||
@@ -201,7 +238,7 @@ async fn publish_status(
|
||||
async fn handle_task(
|
||||
client: &AsyncClient,
|
||||
args: &Args,
|
||||
config: &Arc<RwLock<Option<ProbeConfig>>>,
|
||||
config: &Arc<RwLock<Option<LoadedProbeConfig>>>,
|
||||
topic: &str,
|
||||
task: ProbeTask<'_>,
|
||||
received_at: Instant,
|
||||
@@ -219,180 +256,52 @@ async fn handle_task(
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let config = config.read().await.clone();
|
||||
let Some(config) = config else {
|
||||
bail!("no config");
|
||||
};
|
||||
let result_topic = format!("probe/results/v1/{job_id}/{}", args.probe_id);
|
||||
let probing = join_all(config.hosts.into_iter().map(|host| {
|
||||
let target = task.target.to_string();
|
||||
async move {
|
||||
let probe_evidence = probe_host(&host, &target).await;
|
||||
HostProbeResult {
|
||||
probe_evidence,
|
||||
host_id: host.id,
|
||||
}
|
||||
let config = config.read().await.clone();
|
||||
let control_target = config.as_ref().and_then(|config| {
|
||||
let mut rng = rand::thread_rng();
|
||||
match task.ip {
|
||||
IpAddr::V4(_) => config
|
||||
.control_hosts_v4
|
||||
.choose(&mut rng)
|
||||
.copied()
|
||||
.map(IpAddr::V4),
|
||||
IpAddr::V6(_) => config
|
||||
.control_hosts_v6
|
||||
.choose(&mut rng)
|
||||
.copied()
|
||||
.map(IpAddr::V6),
|
||||
}
|
||||
}));
|
||||
let result = match time::timeout(remaining, probing).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(
|
||||
"dropping expired task {job_id}: timeout {}ms",
|
||||
task.timeout_ms
|
||||
);
|
||||
return Ok(());
|
||||
});
|
||||
let sni_check = sni::check_sni(
|
||||
config.as_ref().map(|config| &config.config),
|
||||
task.domain,
|
||||
remaining,
|
||||
job_id,
|
||||
task.timeout_ms,
|
||||
);
|
||||
let target_traceroute = traceroute::tcp_traceroute(task.ip, args.traceroute_max_hops);
|
||||
let control_traceroute = async {
|
||||
match control_target {
|
||||
Some(target) => traceroute::tcp_traceroute(target, args.traceroute_max_hops).await,
|
||||
None => None,
|
||||
}
|
||||
};
|
||||
let (responses, target_traceroute, control_traceroute) =
|
||||
tokio::join!(sni_check, target_traceroute, control_traceroute);
|
||||
let responses = responses?;
|
||||
|
||||
client
|
||||
.publish(
|
||||
result_topic,
|
||||
QoS::AtLeastOnce,
|
||||
false,
|
||||
serde_json::to_vec(&result)?,
|
||||
serde_json::to_vec(&ProbeResult {
|
||||
responses,
|
||||
target_traceroute,
|
||||
control_traceroute,
|
||||
})?,
|
||||
)
|
||||
.await
|
||||
.context("publish probe result")
|
||||
}
|
||||
|
||||
async fn probe_host(host: &Host, target: &str) -> ProbeEvidence {
|
||||
let timeout = Duration::from_secs(host.timeout_sec as u64);
|
||||
let tcp = match time::timeout(timeout, TcpStream::connect((host.host.as_str(), 443))).await {
|
||||
Ok(Ok(tcp)) => tcp,
|
||||
Ok(Err(_)) | Err(_) => return ProbeEvidence::ConnectionError,
|
||||
};
|
||||
|
||||
let tls_config = ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(Arc::new(NoCertificateVerification))
|
||||
.with_no_client_auth();
|
||||
let connector = TlsConnector::from(Arc::new(tls_config));
|
||||
|
||||
let server_name = match ServerName::try_from(target.to_string()) {
|
||||
Ok(server_name) => server_name,
|
||||
Err(_) => return ProbeEvidence::ClientHello,
|
||||
};
|
||||
|
||||
let mut tls = match time::timeout(timeout, connector.connect(server_name, tcp)).await {
|
||||
Ok(Ok(tls)) => tls,
|
||||
Ok(Err(_)) | Err(_) => return ProbeEvidence::ClientHello,
|
||||
};
|
||||
|
||||
let request = format!(
|
||||
"GET /{} HTTP/1.1\r\nHost: {}\r\nUser-Agent: cheburcheck-probe/{}\r\nRange: bytes=0-{}\r\nConnection: close\r\n\r\n",
|
||||
host.file_path.trim_start_matches('/'),
|
||||
target,
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
host.min_data.saturating_sub(1)
|
||||
);
|
||||
|
||||
if !matches!(
|
||||
time::timeout(timeout, tls.write_all(request.as_bytes())).await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
return ProbeEvidence::ClientHello;
|
||||
}
|
||||
|
||||
let mut received = 0u32;
|
||||
let mut headers_done = false;
|
||||
let mut pending = Vec::new();
|
||||
let mut buffer = [0u8; 8192];
|
||||
loop {
|
||||
match time::timeout(timeout, tls.read(&mut buffer)).await {
|
||||
Ok(Ok(0)) | Err(_) => {
|
||||
return if received >= host.min_data {
|
||||
ProbeEvidence::Good
|
||||
} else {
|
||||
ProbeEvidence::DataTimeout { bytes: received }
|
||||
};
|
||||
}
|
||||
Ok(Ok(bytes)) => {
|
||||
add_response_body_bytes(
|
||||
&buffer[..bytes],
|
||||
&mut pending,
|
||||
&mut headers_done,
|
||||
&mut received,
|
||||
);
|
||||
if received >= host.min_data {
|
||||
return ProbeEvidence::Good;
|
||||
}
|
||||
}
|
||||
Ok(Err(_)) => {
|
||||
return if received >= host.min_data {
|
||||
ProbeEvidence::Good
|
||||
} else {
|
||||
ProbeEvidence::DataTimeout { bytes: received }
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn add_response_body_bytes(
|
||||
chunk: &[u8],
|
||||
pending: &mut Vec<u8>,
|
||||
headers_done: &mut bool,
|
||||
received: &mut u32,
|
||||
) {
|
||||
if *headers_done {
|
||||
*received = received.saturating_add(chunk.len() as u32);
|
||||
return;
|
||||
}
|
||||
|
||||
pending.extend_from_slice(chunk);
|
||||
if let Some(body_start) = pending.windows(4).position(|window| window == b"\r\n\r\n") {
|
||||
*headers_done = true;
|
||||
let body_bytes = pending.len().saturating_sub(body_start + 4);
|
||||
*received = received.saturating_add(body_bytes as u32);
|
||||
pending.clear();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct NoCertificateVerification;
|
||||
|
||||
impl ServerCertVerifier for NoCertificateVerification {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
_end_entity: &CertificateDer<'_>,
|
||||
_intermediates: &[CertificateDer<'_>],
|
||||
_server_name: &ServerName<'_>,
|
||||
_ocsp_response: &[u8],
|
||||
_now: UnixTime,
|
||||
) -> Result<ServerCertVerified, TlsError> {
|
||||
Ok(ServerCertVerified::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, TlsError> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, TlsError> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
|
||||
vec![
|
||||
SignatureScheme::ECDSA_NISTP256_SHA256,
|
||||
SignatureScheme::ECDSA_NISTP384_SHA384,
|
||||
SignatureScheme::ED25519,
|
||||
SignatureScheme::RSA_PSS_SHA256,
|
||||
SignatureScheme::RSA_PSS_SHA384,
|
||||
SignatureScheme::RSA_PSS_SHA512,
|
||||
SignatureScheme::RSA_PKCS1_SHA256,
|
||||
SignatureScheme::RSA_PKCS1_SHA384,
|
||||
SignatureScheme::RSA_PKCS1_SHA512,
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
use anyhow::{Result, bail};
|
||||
use futures::future::join_all;
|
||||
use log::warn;
|
||||
use reports::probe::{Host, HostProbeResult, ProbeConfig, ProbeEvidence};
|
||||
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
|
||||
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
|
||||
use rustls::{ClientConfig, DigitallySignedStruct, Error as TlsError, SignatureScheme};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio::time;
|
||||
use tokio_rustls::TlsConnector;
|
||||
|
||||
pub async fn check_sni(
|
||||
config: Option<&ProbeConfig>,
|
||||
domain: Option<&str>,
|
||||
timeout: Duration,
|
||||
job_id: &str,
|
||||
task_timeout_ms: u64,
|
||||
) -> Result<Option<Vec<HostProbeResult>>> {
|
||||
let Some(domain) = domain else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(config) = config else {
|
||||
bail!("no config");
|
||||
};
|
||||
let probing = join_all(config.hosts.iter().map(|host| async move {
|
||||
let probe_evidence = probe_host(host, domain).await;
|
||||
HostProbeResult {
|
||||
probe_evidence,
|
||||
host_id: host.id.clone(),
|
||||
}
|
||||
}));
|
||||
|
||||
match time::timeout(timeout, probing).await {
|
||||
Ok(responses) => Ok(Some(responses)),
|
||||
Err(_) => {
|
||||
warn!("SNI checks expired for task {job_id}: timeout {task_timeout_ms}ms");
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn probe_host(host: &Host, target: &str) -> ProbeEvidence {
|
||||
let timeout = Duration::from_secs(host.timeout_sec as u64);
|
||||
let tcp = match time::timeout(timeout, TcpStream::connect((host.host.as_str(), 443))).await {
|
||||
Ok(Ok(tcp)) => tcp,
|
||||
Ok(Err(_)) | Err(_) => return ProbeEvidence::ConnectionError,
|
||||
};
|
||||
|
||||
let tls_config = ClientConfig::builder()
|
||||
.dangerous()
|
||||
.with_custom_certificate_verifier(Arc::new(NoCertificateVerification))
|
||||
.with_no_client_auth();
|
||||
let connector = TlsConnector::from(Arc::new(tls_config));
|
||||
|
||||
let server_name = match ServerName::try_from(target.to_string()) {
|
||||
Ok(server_name) => server_name,
|
||||
Err(_) => return ProbeEvidence::ClientHello,
|
||||
};
|
||||
|
||||
let mut tls = match time::timeout(timeout, connector.connect(server_name, tcp)).await {
|
||||
Ok(Ok(tls)) => tls,
|
||||
Ok(Err(_)) | Err(_) => return ProbeEvidence::ClientHello,
|
||||
};
|
||||
|
||||
let request = format!(
|
||||
"GET /{} HTTP/1.1\r\nHost: {}\r\nUser-Agent: cheburcheck-probe/{}\r\nRange: bytes=0-{}\r\nConnection: close\r\n\r\n",
|
||||
host.file_path.trim_start_matches('/'),
|
||||
target,
|
||||
env!("CARGO_PKG_VERSION"),
|
||||
host.min_data.saturating_sub(1)
|
||||
);
|
||||
|
||||
if !matches!(
|
||||
time::timeout(timeout, tls.write_all(request.as_bytes())).await,
|
||||
Ok(Ok(()))
|
||||
) {
|
||||
return ProbeEvidence::ClientHello;
|
||||
}
|
||||
|
||||
let mut received = 0u32;
|
||||
let mut headers_done = false;
|
||||
let mut pending = Vec::new();
|
||||
let mut buffer = [0u8; 8192];
|
||||
loop {
|
||||
match time::timeout(timeout, tls.read(&mut buffer)).await {
|
||||
Ok(Ok(0)) | Err(_) => {
|
||||
return if received >= host.min_data {
|
||||
ProbeEvidence::Good
|
||||
} else {
|
||||
ProbeEvidence::DataTimeout { bytes: received }
|
||||
};
|
||||
}
|
||||
Ok(Ok(bytes)) => {
|
||||
add_response_body_bytes(
|
||||
&buffer[..bytes],
|
||||
&mut pending,
|
||||
&mut headers_done,
|
||||
&mut received,
|
||||
);
|
||||
if received >= host.min_data {
|
||||
return ProbeEvidence::Good;
|
||||
}
|
||||
}
|
||||
Ok(Err(_)) => {
|
||||
return if received >= host.min_data {
|
||||
ProbeEvidence::Good
|
||||
} else {
|
||||
ProbeEvidence::DataTimeout { bytes: received }
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn add_response_body_bytes(
|
||||
chunk: &[u8],
|
||||
pending: &mut Vec<u8>,
|
||||
headers_done: &mut bool,
|
||||
received: &mut u32,
|
||||
) {
|
||||
if *headers_done {
|
||||
*received = received.saturating_add(chunk.len() as u32);
|
||||
return;
|
||||
}
|
||||
|
||||
pending.extend_from_slice(chunk);
|
||||
if let Some(body_start) = pending.windows(4).position(|window| window == b"\r\n\r\n") {
|
||||
*headers_done = true;
|
||||
let body_bytes = pending.len().saturating_sub(body_start + 4);
|
||||
*received = received.saturating_add(body_bytes as u32);
|
||||
pending.clear();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct NoCertificateVerification;
|
||||
|
||||
impl ServerCertVerifier for NoCertificateVerification {
|
||||
fn verify_server_cert(
|
||||
&self,
|
||||
_end_entity: &CertificateDer<'_>,
|
||||
_intermediates: &[CertificateDer<'_>],
|
||||
_server_name: &ServerName<'_>,
|
||||
_ocsp_response: &[u8],
|
||||
_now: UnixTime,
|
||||
) -> Result<ServerCertVerified, TlsError> {
|
||||
Ok(ServerCertVerified::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls12_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, TlsError> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn verify_tls13_signature(
|
||||
&self,
|
||||
_message: &[u8],
|
||||
_cert: &CertificateDer<'_>,
|
||||
_dss: &DigitallySignedStruct,
|
||||
) -> Result<HandshakeSignatureValid, TlsError> {
|
||||
Ok(HandshakeSignatureValid::assertion())
|
||||
}
|
||||
|
||||
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
|
||||
vec![
|
||||
SignatureScheme::ECDSA_NISTP256_SHA256,
|
||||
SignatureScheme::ECDSA_NISTP384_SHA384,
|
||||
SignatureScheme::ED25519,
|
||||
SignatureScheme::RSA_PSS_SHA256,
|
||||
SignatureScheme::RSA_PSS_SHA384,
|
||||
SignatureScheme::RSA_PSS_SHA512,
|
||||
SignatureScheme::RSA_PKCS1_SHA256,
|
||||
SignatureScheme::RSA_PKCS1_SHA384,
|
||||
SignatureScheme::RSA_PKCS1_SHA512,
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,286 @@
|
||||
use etherparse::{
|
||||
Icmpv4Type, Icmpv6Slice, Icmpv6Type, IpNumber, LaxNetSlice, LaxSlicedPacket, TransportSlice,
|
||||
icmpv4, icmpv6,
|
||||
};
|
||||
use polling::{Event, Events, Poller};
|
||||
use reports::probe::{TcpTracerouteOutcome, TcpTracerouteResult};
|
||||
use socket2::{Domain, Protocol, SockAddr, Socket, Type};
|
||||
use std::io::{self, Read};
|
||||
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 tcp_traceroute(target: IpAddr, max_hops: u8) -> Option<TcpTracerouteResult> {
|
||||
tokio::task::spawn_blocking(move || trace_blocking(target, max_hops))
|
||||
.await
|
||||
.map_err(|error| log::warn!("TCP traceroute task failed for {target}: {error}"))
|
||||
.ok()?
|
||||
.map_err(|error| log::warn!("TCP traceroute failed for {target}: {error}"))
|
||||
.ok()
|
||||
}
|
||||
|
||||
fn trace_blocking(target: IpAddr, max_hops: u8) -> io::Result<TcpTracerouteResult> {
|
||||
let (domain, icmp_protocol) = match target {
|
||||
IpAddr::V4(_) => (Domain::IPV4, Protocol::ICMPV4),
|
||||
IpAddr::V6(_) => (Domain::IPV6, Protocol::ICMPV6),
|
||||
};
|
||||
let receiver = Socket::new(domain, Type::RAW, Some(icmp_protocol))?;
|
||||
let mut last_icmp_hop = None;
|
||||
|
||||
for ttl in 1..=max_hops {
|
||||
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(_) => SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
|
||||
IpAddr::V6(_) => SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 0),
|
||||
};
|
||||
tcp.bind(&SockAddr::from(unspecified))?;
|
||||
let destination = SockAddr::from(SocketAddr::new(target, HTTPS_PORT));
|
||||
if let Err(error) = tcp.connect(&destination) {
|
||||
if error.kind() == io::ErrorKind::ConnectionRefused {
|
||||
return Ok(TcpTracerouteResult {
|
||||
target,
|
||||
result: TcpTracerouteOutcome::Rst { hop: ttl },
|
||||
});
|
||||
}
|
||||
if !is_connect_in_progress(&error) {
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
let source_port = tcp
|
||||
.local_addr()?
|
||||
.as_socket()
|
||||
.map(|addr| addr.port())
|
||||
.unwrap_or(0);
|
||||
match wait_for_hop_response(&receiver, &tcp, target, source_port)? {
|
||||
HopResponse::IcmpTimeExceeded => last_icmp_hop = Some(ttl),
|
||||
HopResponse::Rst => {
|
||||
return Ok(TcpTracerouteResult {
|
||||
target,
|
||||
result: TcpTracerouteOutcome::Rst { hop: ttl },
|
||||
});
|
||||
}
|
||||
HopResponse::Connected => {
|
||||
return Ok(TcpTracerouteResult {
|
||||
target,
|
||||
result: TcpTracerouteOutcome::Connected { hop: ttl },
|
||||
});
|
||||
}
|
||||
HopResponse::Timeout => {}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(TcpTracerouteResult {
|
||||
target,
|
||||
result: last_icmp_hop
|
||||
.map(|hop| TcpTracerouteOutcome::IcmpTimeExceeded { hop })
|
||||
.unwrap_or(TcpTracerouteOutcome::Timeout),
|
||||
})
|
||||
}
|
||||
|
||||
fn is_connect_in_progress(error: &io::Error) -> bool {
|
||||
error.kind() == io::ErrorKind::WouldBlock
|
||||
|| error.raw_os_error() == Some(rustix::io::Errno::INPROGRESS.raw_os_error())
|
||||
}
|
||||
|
||||
enum HopResponse {
|
||||
IcmpTimeExceeded,
|
||||
Rst,
|
||||
Connected,
|
||||
Timeout,
|
||||
}
|
||||
|
||||
fn wait_for_hop_response(
|
||||
receiver: &Socket,
|
||||
tcp: &Socket,
|
||||
target: IpAddr,
|
||||
source_port: u16,
|
||||
) -> io::Result<HopResponse> {
|
||||
const ICMP_KEY: usize = 1;
|
||||
const TCP_KEY: usize = 2;
|
||||
|
||||
let poller = Poller::new()?;
|
||||
// SAFETY: both sockets remain alive and are removed from the poller before this function exits.
|
||||
unsafe {
|
||||
poller.add(receiver, Event::readable(ICMP_KEY))?;
|
||||
if let Err(error) = poller.add(tcp, Event::writable(TCP_KEY)) {
|
||||
poller.delete(receiver)?;
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
|
||||
let result = wait_on_poller(&poller, receiver, tcp, target, source_port);
|
||||
let tcp_delete = poller.delete(tcp);
|
||||
let receiver_delete = poller.delete(receiver);
|
||||
let cleanup = tcp_delete.and(receiver_delete);
|
||||
match result {
|
||||
Ok(response) => cleanup.map(|()| response),
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn wait_on_poller(
|
||||
poller: &Poller,
|
||||
receiver: &Socket,
|
||||
tcp: &Socket,
|
||||
target: IpAddr,
|
||||
source_port: u16,
|
||||
) -> io::Result<HopResponse> {
|
||||
const ICMP_KEY: usize = 1;
|
||||
const TCP_KEY: usize = 2;
|
||||
|
||||
let deadline = Instant::now() + HOP_TIMEOUT;
|
||||
let mut watch_tcp = true;
|
||||
let mut buffer = [0u8; 2048];
|
||||
let mut events = Events::new();
|
||||
|
||||
loop {
|
||||
let remaining = deadline.saturating_duration_since(Instant::now());
|
||||
if remaining.is_zero() {
|
||||
return Ok(HopResponse::Timeout);
|
||||
}
|
||||
events.clear();
|
||||
match poller.wait(&mut events, Some(remaining)) {
|
||||
Ok(0) => return Ok(HopResponse::Timeout),
|
||||
Ok(_) => {}
|
||||
Err(error) if error.kind() == io::ErrorKind::Interrupted => continue,
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
|
||||
if watch_tcp && events.iter().any(|event| event.key == TCP_KEY) {
|
||||
match tcp.take_error()? {
|
||||
Some(error) if error.kind() == io::ErrorKind::ConnectionRefused => {
|
||||
return Ok(HopResponse::Rst);
|
||||
}
|
||||
None => return Ok(HopResponse::Connected),
|
||||
Some(_) => watch_tcp = false,
|
||||
}
|
||||
}
|
||||
|
||||
if events.iter().any(|event| event.key == ICMP_KEY) {
|
||||
let mut raw = receiver;
|
||||
match raw.read(&mut buffer) {
|
||||
Ok(bytes) if is_matching_time_exceeded(&buffer[..bytes], target, source_port) => {
|
||||
return Ok(HopResponse::IcmpTimeExceeded);
|
||||
}
|
||||
Ok(_) => poller.modify(receiver, Event::readable(ICMP_KEY))?,
|
||||
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
|
||||
poller.modify(receiver, Event::readable(ICMP_KEY))?;
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
fn is_matching_time_exceeded(packet: &[u8], target: IpAddr, source_port: u16) -> bool {
|
||||
match target {
|
||||
IpAddr::V4(target) => matching_ipv4_time_exceeded(packet, target, source_port),
|
||||
IpAddr::V6(target) => matching_ipv6_time_exceeded(packet, target, source_port),
|
||||
}
|
||||
}
|
||||
|
||||
fn matching_ipv4_time_exceeded(packet: &[u8], target: Ipv4Addr, source_port: u16) -> bool {
|
||||
let Ok(outer) = LaxSlicedPacket::from_ip(packet) else {
|
||||
return false;
|
||||
};
|
||||
let Some(TransportSlice::Icmpv4(icmp)) = outer.transport else {
|
||||
return false;
|
||||
};
|
||||
matches!(
|
||||
icmp.icmp_type(),
|
||||
Icmpv4Type::TimeExceeded(icmpv4::TimeExceededCode::TtlExceededInTransit)
|
||||
) && matching_quoted_tcp(icmp.payload(), IpAddr::V4(target), source_port)
|
||||
}
|
||||
|
||||
fn matching_ipv6_time_exceeded(packet: &[u8], target: Ipv6Addr, source_port: u16) -> bool {
|
||||
let icmp = if packet.first().map(|byte| byte >> 4) == Some(6) {
|
||||
let Ok(outer) = LaxSlicedPacket::from_ip(packet) else {
|
||||
return false;
|
||||
};
|
||||
let Some(TransportSlice::Icmpv6(icmp)) = outer.transport else {
|
||||
return false;
|
||||
};
|
||||
icmp
|
||||
} else {
|
||||
let Ok(icmp) = Icmpv6Slice::from_slice(packet) else {
|
||||
return false;
|
||||
};
|
||||
icmp
|
||||
};
|
||||
matches!(
|
||||
icmp.icmp_type(),
|
||||
Icmpv6Type::TimeExceeded(icmpv6::TimeExceededCode::HopLimitExceeded)
|
||||
) && matching_quoted_tcp(icmp.payload(), IpAddr::V6(target), source_port)
|
||||
}
|
||||
|
||||
fn matching_quoted_tcp(packet: &[u8], target: IpAddr, source_port: u16) -> bool {
|
||||
let Ok(quoted) = LaxSlicedPacket::from_ip(packet) else {
|
||||
return false;
|
||||
};
|
||||
let (ip_number, payload) = match (quoted.net, target) {
|
||||
(Some(LaxNetSlice::Ipv4(ip)), IpAddr::V4(target))
|
||||
if ip.header().destination_addr() == target =>
|
||||
{
|
||||
(ip.payload().ip_number, ip.payload().payload)
|
||||
}
|
||||
(Some(LaxNetSlice::Ipv6(ip)), IpAddr::V6(target))
|
||||
if ip.header().destination_addr() == target =>
|
||||
{
|
||||
(ip.payload().ip_number, ip.payload().payload)
|
||||
}
|
||||
_ => return false,
|
||||
};
|
||||
if ip_number != IpNumber::TCP {
|
||||
return false;
|
||||
}
|
||||
let Some(ports) = payload.get(..4) else {
|
||||
return false;
|
||||
};
|
||||
u16::from_be_bytes([ports[0], ports[1]]) == source_port
|
||||
&& u16::from_be_bytes([ports[2], ports[3]]) == HTTPS_PORT
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn matches_ipv4_time_exceeded_for_tcp_flow() {
|
||||
let target = Ipv4Addr::new(203, 0, 113, 10);
|
||||
let mut packet = vec![0; 20 + 8 + 20 + 8];
|
||||
packet[0] = 0x45;
|
||||
packet[2..4].copy_from_slice(&56u16.to_be_bytes());
|
||||
packet[9] = IpNumber::ICMP.0;
|
||||
packet[20] = 11;
|
||||
packet[28] = 0x45;
|
||||
packet[30..32].copy_from_slice(&28u16.to_be_bytes());
|
||||
packet[37] = 6;
|
||||
packet[44..48].copy_from_slice(&target.octets());
|
||||
packet[48..50].copy_from_slice(&42_000u16.to_be_bytes());
|
||||
packet[50..52].copy_from_slice(&HTTPS_PORT.to_be_bytes());
|
||||
|
||||
assert!(matching_ipv4_time_exceeded(&packet, target, 42_000));
|
||||
assert!(!matching_ipv4_time_exceeded(&packet, target, 42_001));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn matches_ipv6_time_exceeded_for_tcp_flow() {
|
||||
let target = "2001:db8::10".parse::<Ipv6Addr>().unwrap();
|
||||
let mut packet = vec![0; 8 + 40 + 8];
|
||||
packet[0] = 3;
|
||||
packet[8] = 0x60;
|
||||
packet[14] = 6;
|
||||
packet[32..48].copy_from_slice(&target.octets());
|
||||
packet[48..50].copy_from_slice(&42_000u16.to_be_bytes());
|
||||
packet[50..52].copy_from_slice(&HTTPS_PORT.to_be_bytes());
|
||||
|
||||
assert!(matching_ipv6_time_exceeded(&packet, target, 42_000));
|
||||
assert!(!matching_ipv6_time_exceeded(&packet, target, 42_001));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user