diff --git a/Cargo.lock b/Cargo.lock index 59d33b4..62686ed 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1223,7 +1223,7 @@ dependencies = [ "idna", "ipnet", "jni", - "rand 0.10.2", + "rand 0.10.3", "rustls 0.23.44", "thiserror 2.0.20", "tinyvec", @@ -1246,7 +1246,7 @@ dependencies = [ "jni", "once_cell", "prefix-trie 0.8.4", - "rand 0.10.2", + "rand 0.10.3", "ring", "thiserror 2.0.20", "tinyvec", @@ -1271,7 +1271,7 @@ dependencies = [ "ndk-context", "once_cell", "parking_lot", - "rand 0.10.2", + "rand 0.10.3", "resolv-conf", "rustls 0.23.44", "smallvec", @@ -2603,7 +2603,7 @@ dependencies = [ "bytes", "getrandom 0.4.3", "lru-slab", - "rand 0.10.2", + "rand 0.10.3", "rand_pcg", "ring", "rustc-hash", @@ -2674,9 +2674,9 @@ dependencies = [ [[package]] name = "rand" -version = "0.10.2" +version = "0.10.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +checksum = "65c9fb96cbc91e3478eaae79a69fcd3f1ae4ad052e471fe6732fff548984b4af" dependencies = [ "chacha20", "getrandom 0.4.3", @@ -4517,6 +4517,7 @@ dependencies = [ "governor", "log", "querying", + "rand 0.10.3", "reports", "reqwest", "rocket", @@ -4528,6 +4529,7 @@ dependencies = [ "serde", "sqlx", "toml", + "uuid", ] [[package]] diff --git a/README.md b/README.md index 40e23e3..0c9d69c 100644 --- a/README.md +++ b/README.md @@ -72,6 +72,21 @@ HTTP_PORT=80 docker compose up --build `DATABASE_INTERVAL_SECONDS`, а интервал повтора после ошибки скачивания — через `DATABASE_RETRY_INTERVAL_SECONDS` (по умолчанию 300 секунд). +## Управление сканерами + +Задайте `MQTT_ADMIN_PASSWORD` в `.env` и откройте `/admin/probes`. Пароль защищает +административный API и после ввода сохраняется в localStorage браузера. Кнопка +«Выйти» удаляет его из localStorage. Для доставки команд также требуется +`MQTT_ADMIN_TOKEN` — учётные данные веб-сервера для MQTT. + +На странице можно создавать сканеры, менять их метаданные и флаги, отправлять +команды переподписки и traceroute, а также запускать проверку обновлений для +всех подключённых сканеров или одного сканера. `PROBE_ID` и `PROBE_TOKEN` нового сканера +показываются сразу после создания; сохраните токен, так как повторно он не +отображается. Конфигурация `website/probe-hosts.toml` монтируется в контейнер +веб-сервера: кнопка перезагрузки читает файл с диска и публикует новую +конфигурацию в MQTT. + --- ## Вклад diff --git a/docker-compose.yaml b/docker-compose.yaml index 2ec0d01..9ecdad3 100644 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -32,6 +32,8 @@ services: DATABASE_CACHE_DIR: /var/cache/cheburcheck/databases DATABASE_RETRY_INTERVAL_SECONDS: "${DATABASE_RETRY_INTERVAL_SECONDS:-300}" MQTT_ADMIN_TOKEN: "${MQTT_ADMIN_TOKEN}" + MQTT_ADMIN_PASSWORD: "${MQTT_ADMIN_PASSWORD}" + PROBE_CONFIG_PATH: /app/probe-hosts.toml NODE_STATS_KEY: "${NODE_STATS_KEY}" MQTT_HOST: rmqtt MQTT_PORT: 11883 @@ -46,6 +48,7 @@ services: GITHUB_TOKEN: "${GITHUB_TOKEN:-}" volumes: - database-cache:/var/cache/cheburcheck/databases + - ./website/probe-hosts.toml:/app/probe-hosts.toml:ro expose: - "8000" diff --git a/frontend/src/lib/api/admin.ts b/frontend/src/lib/api/admin.ts new file mode 100644 index 0000000..423aa4e --- /dev/null +++ b/frontend/src/lib/api/admin.ts @@ -0,0 +1,24 @@ +export async function adminRequest( + password: string, + path: string, + method = "GET", + body?: unknown, +): Promise { + const response = await fetch("/api/v1/admin" + path, { + method, + cache: "no-store", + headers: { + Authorization: "Bearer " + password, + Accept: "application/json", + ...(body ? { "Content-Type": "application/json" } : {}), + }, + ...(body ? { body: JSON.stringify(body) } : {}), + }); + if (!response.ok) + throw new Error( + response.status === 401 + ? "Неверный пароль администратора" + : "Ошибка " + response.status + ": " + response.statusText, + ); + return response.json() as Promise; +} diff --git a/frontend/src/routes/admin/probes/+page.svelte b/frontend/src/routes/admin/probes/+page.svelte new file mode 100644 index 0000000..1489b4e --- /dev/null +++ b/frontend/src/routes/admin/probes/+page.svelte @@ -0,0 +1,853 @@ + + +Сканеры · Cheburcheck + + +
+
+

Сканеры

+ {#if password} + + {/if} +
+ + {#if !password || error === "Неверный пароль администратора"} +
{ event.preventDefault(); void signIn(); }} + > + + + {#if error} + + {/if} + +
+ {:else} +
+ + +
+ {#if notice} +
+ {notice} +
+ {/if} + {#if error} + + {/if} + +
+
+

Текущие сканеры

+
+ + {probes.length} + всего · {probes.filter((p) => p.online).length} онлайн +
+
+ {#if visibleProbes.length === 0} +

+ {loading ? "Загрузка…" : onlineOnly && probes.length > 0 ? "Нет сканеров онлайн" : "Сканеров пока нет"} +

+ {:else} + + + + + + + + + + + + {#each visibleProbes as probe (probe.id)} + + + + + + + + {/each} + +
СканерСостояниеМетаданныеФлагиДействия
+
#{probe.id}
+
+

{probe.name}

+ {probe.last_connected_at ? `${new Date(probe.last_connected_at).toLocaleString("ru-RU")}` : "Не подключался"} +
+
+ + {probe.online ? "Онлайн" : "Офлайн"} +
+ {probe.version ?? "Версия неизвестна"} + {probe.bundle_type ? `· ${probe.bundle_type}` : ""} +
+ {#if probe.dpi_hop_v4 || probe.dpi_hop_v6} +
+ v4: {probe.dpi_hop_v4 ?? "—"}; v6: + {probe.dpi_hop_v6 ?? "—"} +
+ {/if} +
+ {probe.region || "Регион не указан"}
+ {probe.provider || "Провайдер не указан"} + {#if probe.asn} + · + {probe.asn} + {/if} +
+
+ +
+
+
+ + +
+
+ {/if} +
+ + {#if created} +
+

Данные нового сканера

+

+ Токен показан только сейчас. Скопируйте его перед закрытием страницы. +

+
+
PROBE_ID={created.id}
+PROBE_TOKEN={created.token}
+ +
+
+ {/if} + +
+

+ {editing === null ? "Добавить сканер" : `Изменить сканер #${editing}`} +

+
{ event.preventDefault(); void save(); }} + > +
+ +
+
+ +
+
+ + {#if editing !== null} + + {/if} +
+
+
+ +
+

Команды и трассировка

+

+ Команда выполняется выбранным сканером. Ответ может занять до минуты. +

+
+ +
+ {#if busy.endsWith(":command")} +

Ожидание ответа сканера…

+ {/if} + {#if result?.type === "error"} +

{result.message}

+ {:else if result?.type === "resubscribe_tasks"} +

+ Запрос на переподписку получен сканером. +

+ {:else if result?.type === "traceroute"} +
+

Маршрут до {result.target}

+
    + {#each result.hops as hop} +
  1. + {hop.ttl} +
    + {hop.address ?? "* * *"} + {#if hop.reverse_names.length} + {hop.reverse_names.join(", ")} + {/if} +
    + {outcomeLabel[hop.outcome] ?? hop.outcome} +
  2. + {/each} +
+ {#if result.hops.length === 0} +

Ответов нет.

+ {/if} +
+ {/if} +
+ {/if} +
+ + diff --git a/nginx.conf b/nginx.conf index e9c61d8..b92cfba 100644 --- a/nginx.conf +++ b/nginx.conf @@ -32,7 +32,7 @@ server { location /api/v1/ { proxy_buffering off; proxy_cache off; - proxy_read_timeout 60s; + proxy_read_timeout 75s; proxy_pass http://website_backend; } diff --git a/website/Cargo.toml b/website/Cargo.toml index 077e2d8..412d953 100644 --- a/website/Cargo.toml +++ b/website/Cargo.toml @@ -20,4 +20,6 @@ dotenvy = { version = "0.15.7" } governor = { version = "0.6", features = ["dashmap"] } rumqttc = "0.24" toml = "0.8" +uuid = { version = "1", features = ["v4"] } reqwest = { workspace = true } +rand = "0.10.3" diff --git a/website/migrations/20260926000001_reports_reporter_nullable.sql b/website/migrations/20260926000001_reports_reporter_nullable.sql new file mode 100644 index 0000000..a01b2bb --- /dev/null +++ b/website/migrations/20260926000001_reports_reporter_nullable.sql @@ -0,0 +1,2 @@ +-- Keep historical reports when their reporter is removed. +ALTER TABLE reports ALTER COLUMN reporter DROP NOT NULL; diff --git a/website/src/admin.rs b/website/src/admin.rs new file mode 100644 index 0000000..65c2e18 --- /dev/null +++ b/website/src/admin.rs @@ -0,0 +1,340 @@ +use crate::mqtt::{MqttPublisher, ProbeStatusSnapshot}; +use rand::distr::{Alphanumeric, SampleString}; +use reports::probe::{ProbeCommand, ProbeCommandResult}; +use rocket::http::Status; +use rocket::request::{FromRequest, Outcome}; +use rocket::serde::json::Json; +use rocket::{Request, State}; +use serde::{Deserialize, Serialize}; +use sqlx::PgPool; +use sqlx::types::chrono::{DateTime, Utc}; +use std::net::IpAddr; + +pub struct Admin; + +#[rocket::async_trait] +impl<'r> FromRequest<'r> for Admin { + type Error = (); + + async fn from_request(request: &'r Request<'_>) -> Outcome { + let configured = std::env::var("MQTT_ADMIN_PASSWORD") + .ok() + .filter(|s| !s.is_empty()); + let supplied = request + .headers() + .get_one("Authorization") + .and_then(|s| s.strip_prefix("Bearer ")); + match (configured, supplied) { + (Some(expected), Some(actual)) + if constant_time_eq(expected.as_bytes(), actual.as_bytes()) => + { + Outcome::Success(Admin) + } + _ => Outcome::Error((Status::Unauthorized, ())), + } + } +} + +fn constant_time_eq(a: &[u8], b: &[u8]) -> bool { + let mut difference = a.len() ^ b.len(); + for (left, right) in a.iter().zip(b.iter()) { + difference |= (left ^ right) as usize; + } + difference == 0 +} + +#[derive(Serialize, sqlx::FromRow)] +pub struct ProbeRow { + id: i32, + name: String, + region: Option, + asn: Option, + provider: Option, + hidden: bool, + disable_traceroutes: bool, + cdn_unblocked: bool, + last_connected_at: Option>, + online: bool, + version: Option, + bundle_type: Option, + dpi_hop_v4: Option, + dpi_hop_v6: Option, +} + +async fn rows(pool: &PgPool, mqtt: &MqttPublisher) -> Result, Status> { + let mut rows: Vec = sqlx::query_as( + "SELECT id, name, region, asn, provider, hidden, disable_traceroutes, cdn_unblocked, last_connected_at, + false AS online, NULL::text AS version, NULL::text AS bundle_type, NULL::smallint AS dpi_hop_v4, NULL::smallint AS dpi_hop_v6 + FROM reporters ORDER BY id" + ).fetch_all(pool).await.map_err(|error| { + log::error!("failed to load admin probe list: {error}"); + Status::InternalServerError + })?; + let statuses = mqtt.probe_statuses().await; + for row in &mut rows { + if let Some(ProbeStatusSnapshot { + online, + version, + bundle_type, + dpi_hop_v4, + dpi_hop_v6, + }) = statuses.get(&row.id.to_string()) + { + row.online = *online; + row.version = Some(version.clone()); + row.bundle_type = bundle_type.clone(); + row.dpi_hop_v4 = dpi_hop_v4.map(i16::from); + row.dpi_hop_v6 = dpi_hop_v6.map(i16::from); + } + } + Ok(rows) +} + +#[get("/probes")] +pub async fn list( + _admin: Admin, + pool: &State, + mqtt: &State, +) -> Result>, Status> { + Ok(Json(rows(pool, mqtt).await?)) +} + +#[derive(Deserialize)] +pub struct ProbeInput { + name: String, + region: Option, + asn: Option, + provider: Option, + #[serde(default)] + hidden: bool, + #[serde(default)] + disable_traceroutes: bool, + #[serde(default)] + cdn_unblocked: bool, +} + +fn validate(input: &ProbeInput) -> Result<(), Status> { + if input.name.trim().is_empty() + || input.name.len() > 255 + || input.region.as_ref().is_some_and(|s| s.len() > 255) + || input.asn.as_ref().is_some_and(|s| s.len() > 32) + || input.provider.as_ref().is_some_and(|s| s.len() > 255) + { + return Err(Status::BadRequest); + } + Ok(()) +} + +#[derive(Serialize)] +pub struct CreatedProbe { + id: i32, + token: String, +} + +#[post("/probes", format = "json", data = "")] +pub async fn create( + _admin: Admin, + input: Json, + pool: &State, +) -> Result, Status> { + validate(&input)?; + let token = Alphanumeric.sample_string(&mut rand::rng(), 16); + let id: i32 = sqlx::query_scalar("INSERT INTO reporters (name, token, region, asn, provider, hidden, disable_traceroutes, cdn_unblocked) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) RETURNING id") + .bind(input.name.trim()).bind(&token).bind(&input.region).bind(&input.asn).bind(&input.provider) + .bind(input.hidden).bind(input.disable_traceroutes).bind(input.cdn_unblocked) + .fetch_one(&**pool).await.map_err(|_| Status::InternalServerError)?; + Ok(Json(CreatedProbe { id, token })) +} + +#[put("/probes/", format = "json", data = "")] +pub async fn update( + _admin: Admin, + id: i32, + input: Json, + pool: &State, + mqtt: &State, +) -> Result, Status> { + validate(&input)?; + let updated: Option = sqlx::query_scalar("UPDATE reporters SET name=$2, region=$3, asn=$4, provider=$5, hidden=$6, disable_traceroutes=$7, cdn_unblocked=$8 WHERE id=$1 RETURNING id") + .bind(id).bind(input.name.trim()).bind(&input.region).bind(&input.asn).bind(&input.provider) + .bind(input.hidden).bind(input.disable_traceroutes).bind(input.cdn_unblocked) + .fetch_optional(&**pool).await.map_err(|_| Status::InternalServerError)?; + if updated.is_none() { + return Err(Status::NotFound); + } + rows(pool, mqtt) + .await? + .into_iter() + .find(|row| row.id == id) + .map(Json) + .ok_or(Status::NotFound) +} + +#[derive(Serialize)] +pub struct RemovedProbe { + removed: bool, +} + +#[delete("/probes/")] +pub async fn remove( + _admin: Admin, + id: i32, + pool: &State, +) -> Result, Status> { + let mut tx = pool.begin().await.map_err(|error| { + log::error!("failed to begin probe removal {id}: {error}"); + Status::InternalServerError + })?; + let exists: Option = + sqlx::query_scalar("SELECT id FROM reporters WHERE id = $1 FOR UPDATE") + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_err(|error| { + log::error!("failed to lock probe {id} for removal: {error}"); + Status::InternalServerError + })?; + if exists.is_none() { + return Err(Status::NotFound); + } + sqlx::query("UPDATE reports SET reporter = NULL WHERE reporter = $1") + .bind(id) + .execute(&mut *tx) + .await + .map_err(|error| { + log::error!("failed to detach reports from probe {id}: {error}"); + Status::InternalServerError + })?; + sqlx::query("DELETE FROM reporters WHERE id = $1") + .bind(id) + .execute(&mut *tx) + .await + .map_err(|error| { + log::error!("failed to delete probe {id}: {error}"); + Status::InternalServerError + })?; + tx.commit().await.map_err(|error| { + log::error!("failed to commit probe removal {id}: {error}"); + Status::InternalServerError + })?; + Ok(Json(RemovedProbe { removed: true })) +} + +#[post("/probe-config/reload")] +pub async fn reload( + _admin: Admin, + mqtt: &State, +) -> Result, Status> { + mqtt.reload_probe_config() + .await + .map(Json) + .map_err(|_| Status::InternalServerError) +} + +#[derive(Serialize)] +pub struct UpdateCheckResponse { + requested: bool, +} + +#[post("/probes/update-check")] +pub async fn update_all_probes( + _admin: Admin, + mqtt: &State, +) -> Result, Status> { + mqtt.request_probe_update(None) + .await + .map_err(|_| Status::ServiceUnavailable)?; + Ok(Json(UpdateCheckResponse { requested: true })) +} + +#[post("/probes//update-check")] +pub async fn update_one_probe( + _admin: Admin, + id: i32, + pool: &State, + mqtt: &State, +) -> Result, Status> { + let exists: bool = sqlx::query_scalar("SELECT EXISTS(SELECT 1 FROM reporters WHERE id=$1)") + .bind(id) + .fetch_one(&**pool) + .await + .map_err(|_| Status::InternalServerError)?; + if !exists { + return Err(Status::NotFound); + } + mqtt.request_probe_update(Some(&id.to_string())) + .await + .map_err(|_| Status::ServiceUnavailable)?; + Ok(Json(UpdateCheckResponse { requested: true })) +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CommandInput { + ResubscribeTasks, + Traceroute { target: IpAddr, max_hops: u8 }, +} + +#[post("/probes//commands", format = "json", data = "")] +pub async fn command( + _admin: Admin, + id: i32, + input: Json, + pool: &State, + mqtt: &State, +) -> Result, Status> { + let exists: bool = sqlx::query_scalar("SELECT EXISTS(SELECT 1 FROM reporters WHERE id=$1)") + .bind(id) + .fetch_one(&**pool) + .await + .map_err(|_| Status::InternalServerError)?; + if !exists { + return Err(Status::NotFound); + } + let command = match input.into_inner() { + CommandInput::ResubscribeTasks => ProbeCommand::ResubscribeTasks, + CommandInput::Traceroute { target, max_hops } if (1..=64).contains(&max_hops) => { + ProbeCommand::Traceroute { target, max_hops } + } + CommandInput::Traceroute { .. } => return Err(Status::BadRequest), + }; + mqtt.send_command(&id.to_string(), command) + .await + .map(Json) + .map_err(|error| match error { + crate::mqtt::PublishError::CommandTimeout => Status::GatewayTimeout, + _ => Status::ServiceUnavailable, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn validates_probe_fields() { + let mut input = ProbeInput { + name: "test".into(), + region: None, + asn: None, + provider: None, + hidden: false, + disable_traceroutes: false, + cdn_unblocked: false, + }; + assert!(validate(&input).is_ok()); + input.name = " ".into(); + assert!(validate(&input).is_err()); + input.name = "test".into(); + input.asn = Some("x".repeat(33)); + assert!(validate(&input).is_err()); + } + + #[test] + fn compares_passwords_without_accepting_prefixes() { + assert!(constant_time_eq(b"secret", b"secret")); + assert!(!constant_time_eq(b"secret", b"secre")); + assert!(!constant_time_eq(b"secret", b"secret2")); + } +} diff --git a/website/src/api/probe.rs b/website/src/api/probe.rs index 143befa..eab34f6 100644 --- a/website/src/api/probe.rs +++ b/website/src/api/probe.rs @@ -140,7 +140,7 @@ pub async fn probe_query( } }; let is_ip_target = domain.is_none(); - let probe_config = mqtt.probe_config(); + let probe_config = mqtt.probe_config().await; if domain.is_none() && !probe_config.traceroute_enabled { return Err(Status::BadRequest); } @@ -648,6 +648,7 @@ fn publish_error_status(error: PublishError) -> Status { | PublishError::Serialize(_) | PublishError::Subscribe(_) | PublishError::Publish(_) => Status::InternalServerError, + PublishError::CommandTimeout => Status::GatewayTimeout, } } diff --git a/website/src/main.rs b/website/src/main.rs index 567d697..10dff0e 100644 --- a/website/src/main.rs +++ b/website/src/main.rs @@ -1,5 +1,6 @@ #[macro_use] extern crate rocket; +mod admin; mod agency; mod api; mod database_refresh; @@ -121,6 +122,19 @@ async fn rocket() -> _ { ], ) .mount("/agency", routes![agency::upload_report]) + .mount( + "/api/v1/admin", + routes![ + admin::list, + admin::create, + admin::update, + admin::remove, + admin::reload, + admin::update_all_probes, + admin::update_one_probe, + admin::command + ], + ) .mount("/mqtt", routes![mqtt_auth::auth, mqtt_auth::acl]) .mount( "/api/v1/probe-updates", diff --git a/website/src/mqtt.rs b/website/src/mqtt.rs index 843043b..82734e6 100644 --- a/website/src/mqtt.rs +++ b/website/src/mqtt.rs @@ -1,7 +1,8 @@ use log::{info, warn}; use reports::probe::HostType; use reports::probe::{ - DpiProbeConfig, Host, ProbeConfig, ProbeResult, ProbeResultEvent, ProbeStatus, ProbeTask, + DpiProbeConfig, Host, ProbeCommand, ProbeCommandResult, ProbeConfig, ProbeResult, + ProbeResultEvent, ProbeStatus, ProbeTask, }; use rocket::serde::json::serde_json; use rumqttc::{AsyncClient, Event as MqttEvent, Incoming, MqttOptions, QoS}; @@ -16,8 +17,6 @@ use std::time::Duration; const MQTT_MAX_PACKET_SIZE: usize = 1024 * 1024; -const DEFAULT_PROBE_HOSTS: &str = include_str!("../probe-hosts.toml"); - #[derive(Debug)] pub enum PublishError { NotConfigured, @@ -26,6 +25,7 @@ pub enum PublishError { Serialize(serde_json::Error), Subscribe(rumqttc::ClientError), Publish(rumqttc::ClientError), + CommandTimeout, } impl fmt::Display for PublishError { @@ -45,6 +45,7 @@ impl fmt::Display for PublishError { write!(formatter, "failed to subscribe to results: {error}") } PublishError::Publish(error) => write!(formatter, "failed to publish task: {error}"), + PublishError::CommandTimeout => write!(formatter, "probe command timed out"), } } } @@ -54,7 +55,18 @@ pub struct MqttPublisher { client: Option, sessions: Arc>>, probe_statuses: ProbeStatuses, - probe_config: Arc, + probe_config: Arc>, + command_sessions: Arc< + rocket::tokio::sync::Mutex< + HashMap< + String, + ( + String, + rocket::tokio::sync::oneshot::Sender, + ), + >, + >, + >, task_timeout_ms: u64, } @@ -99,20 +111,23 @@ impl MqttPublisher { let sessions = Arc::new(rocket::tokio::sync::RwLock::new(HashMap::new())); let probe_statuses = Arc::new(rocket::tokio::sync::RwLock::new(HashMap::new())); let task_timeout_ms = task_timeout_ms_from_env(); - let probe_config = Arc::new(load_probe_config(task_timeout_ms).unwrap_or_else(|error| { - warn!("failed to load probe config: {error}"); - ProbeConfig { - version: env!("CARGO_PKG_VERSION").to_string(), - task_timeout_ms, - published_at: Utc::now().to_rfc3339(), - hosts: Vec::new(), - traceroute_enabled: false, - dns_samples_per_protocol: reports::probe::default_dns_samples_per_protocol(), - dns_spoofing_provider_threshold: - reports::probe::default_dns_spoofing_provider_threshold(), - dpi_probe: None, - } - })); + let probe_config = Arc::new(rocket::tokio::sync::RwLock::new( + load_probe_config(task_timeout_ms).unwrap_or_else(|error| { + warn!("failed to load probe config: {error}"); + ProbeConfig { + version: env!("CARGO_PKG_VERSION").to_string(), + task_timeout_ms, + published_at: Utc::now().to_rfc3339(), + hosts: Vec::new(), + traceroute_enabled: false, + dns_samples_per_protocol: reports::probe::default_dns_samples_per_protocol(), + dns_spoofing_provider_threshold: + reports::probe::default_dns_spoofing_provider_threshold(), + dpi_probe: None, + } + }), + )); + let command_sessions = Arc::new(rocket::tokio::sync::Mutex::new(HashMap::new())); let admin_token = match std::env::var("MQTT_ADMIN_TOKEN") { Ok(token) if !token.is_empty() => token, _ => { @@ -122,6 +137,7 @@ impl MqttPublisher { sessions, probe_statuses, probe_config, + command_sessions, task_timeout_ms: task_timeout_ms_from_env(), }; } @@ -145,12 +161,26 @@ impl MqttPublisher { let event_probe_statuses = probe_statuses.clone(); let config_client = client.clone(); let event_probe_config = probe_config.clone(); + let event_command_sessions = command_sessions.clone(); rocket::tokio::spawn(async move { loop { match eventloop.poll().await { Ok(MqttEvent::Incoming(Incoming::ConnAck(_))) => { + if let Err(error) = config_client + .subscribe("probe/status/v1/+", QoS::AtLeastOnce) + .await + { + warn!("failed to subscribe to probe status updates: {error}"); + } + if let Err(error) = config_client + .subscribe("probe/command-results/v1/+/+", QoS::AtLeastOnce) + .await + { + warn!("failed to subscribe to command results: {error}"); + } if let Err(error) = - publish_probe_config(&config_client, event_probe_config.as_ref()).await + publish_probe_config(&config_client, &*event_probe_config.read().await) + .await { warn!("failed to publish retained probe config: {error}"); } @@ -164,6 +194,23 @@ impl MqttPublisher { &publish.payload, ) .await; + if let Some((probe_id, command_id)) = + parse_command_result_topic(&publish.topic) + { + if let Ok(result) = + serde_json::from_slice::(&publish.payload) + { + let mut sessions = event_command_sessions.lock().await; + if sessions + .get(command_id) + .is_some_and(|(expected, _)| expected == probe_id) + { + if let Some((_, sender)) = sessions.remove(command_id) { + let _ = sender.send(result); + } + } + } + } } Ok(_) => {} Err(error) => { @@ -174,22 +221,13 @@ impl MqttPublisher { } }); - let status_client = client.clone(); - rocket::tokio::spawn(async move { - if let Err(error) = status_client - .subscribe("probe/status/v1/+", QoS::AtLeastOnce) - .await - { - warn!("failed to subscribe to probe status updates: {error}"); - } - }); - info!("mqtt publisher configured for {host}:{port}"); Self { client: Some(client), sessions, probe_statuses, probe_config, + command_sessions, task_timeout_ms, } } @@ -215,8 +253,58 @@ impl MqttPublisher { self.probe_statuses.read().await.clone() } - pub fn probe_config(&self) -> Arc { - self.probe_config.clone() + pub async fn probe_config(&self) -> ProbeConfig { + self.probe_config.read().await.clone() + } + + pub async fn reload_probe_config(&self) -> Result { + let config = load_probe_config(self.task_timeout_ms)?; + let client = self.client.as_ref().ok_or(PublishError::NotConfigured)?; + publish_probe_config(client, &config).await?; + *self.probe_config.write().await = config.clone(); + Ok(config) + } + + pub async fn request_probe_update(&self, probe_id: Option<&str>) -> Result<(), PublishError> { + let client = self.client.as_ref().ok_or(PublishError::NotConfigured)?; + let topic = probe_update_topic(probe_id); + client + .publish(topic, QoS::AtLeastOnce, false, b"check".as_slice()) + .await + .map_err(PublishError::Publish) + } + + pub async fn send_command( + &self, + probe_id: &str, + command: ProbeCommand, + ) -> Result { + let client = self.client.as_ref().ok_or(PublishError::NotConfigured)?; + let command_id = uuid::Uuid::new_v4().to_string(); + let (sender, receiver) = rocket::tokio::sync::oneshot::channel(); + self.command_sessions + .lock() + .await + .insert(command_id.clone(), (probe_id.to_owned(), sender)); + let payload = serde_json::to_vec(&command).map_err(PublishError::Serialize)?; + let published = client + .publish( + format!("probe/commands/v1/{probe_id}/{command_id}"), + QoS::AtLeastOnce, + false, + payload, + ) + .await; + if let Err(error) = published { + self.command_sessions.lock().await.remove(&command_id); + return Err(PublishError::Publish(error)); + } + let result = rocket::tokio::time::timeout(Duration::from_secs(60), receiver).await; + self.command_sessions.lock().await.remove(&command_id); + result + .ok() + .and_then(Result::ok) + .ok_or(PublishError::CommandTimeout) } pub async fn subscribe_probe_results( @@ -296,12 +384,17 @@ async fn publish_probe_config( } fn load_probe_config(task_timeout_ms: u64) -> Result { - let config = if let Some(path) = std::env::var_os("PROBE_CONFIG_PATH") { - let contents = std::fs::read_to_string(path).map_err(PublishError::Config)?; - parse_probe_hosts(&contents)? - } else { - parse_probe_hosts(DEFAULT_PROBE_HOSTS)? - }; + let path = std::env::var_os("PROBE_CONFIG_PATH") + .map(std::path::PathBuf::from) + .unwrap_or_else(|| { + if std::path::Path::new("website/probe-hosts.toml").exists() { + "website/probe-hosts.toml".into() + } else { + "probe-hosts.toml".into() + } + }); + let contents = std::fs::read_to_string(path).map_err(PublishError::Config)?; + let config = parse_probe_hosts(&contents)?; Ok(ProbeConfig { version: env!("CARGO_PKG_VERSION").to_string(), task_timeout_ms, @@ -415,6 +508,13 @@ async fn dispatch_probe_result( } } +fn probe_update_topic(probe_id: Option<&str>) -> String { + match probe_id { + Some(id) => format!("probe/update/v1/{id}"), + None => "probe/update/v1".to_string(), + } +} + fn parse_probe_status_topic(topic: &str) -> Option<&str> { let mut parts = topic.split('/'); match ( @@ -429,6 +529,28 @@ fn parse_probe_status_topic(topic: &str) -> Option<&str> { } } +fn parse_command_result_topic(topic: &str) -> Option<(&str, &str)> { + let mut parts = topic.split('/'); + match ( + parts.next(), + parts.next(), + parts.next(), + parts.next(), + parts.next(), + parts.next(), + ) { + ( + Some("probe"), + Some("command-results"), + Some("v1"), + Some(probe_id), + Some(command_id), + None, + ) if !probe_id.is_empty() && !command_id.is_empty() => Some((probe_id, command_id)), + _ => None, + } +} + fn parse_probe_result_topic(topic: &str) -> Option<(&str, &str)> { let mut parts = topic.split('/'); match ( @@ -457,6 +579,28 @@ fn task_timeout_ms_from_env() -> u64 { mod tests { use super::*; + #[test] + fn update_check_uses_global_or_targeted_topic() { + assert_eq!(probe_update_topic(None), "probe/update/v1"); + assert_eq!(probe_update_topic(Some("42")), "probe/update/v1/42"); + } + + #[test] + fn command_result_topics_require_exact_probe_and_command() { + assert_eq!( + parse_command_result_topic("probe/command-results/v1/42/abc"), + Some(("42", "abc")) + ); + assert_eq!( + parse_command_result_topic("probe/command-results/v1/42/abc/extra"), + None + ); + assert_eq!( + parse_command_result_topic("probe/command-results/v1/42/"), + None + ); + } + #[rocket::async_test] async fn status_snapshot_keeps_offline_node_metadata() { let statuses = Arc::new(rocket::tokio::sync::RwLock::new(HashMap::new())); diff --git a/website/src/mqtt_auth.rs b/website/src/mqtt_auth.rs index 92f9ebf..57e5c99 100644 --- a/website/src/mqtt_auth.rs +++ b/website/src/mqtt_auth.rs @@ -109,6 +109,20 @@ pub async fn acl( if request.username != "probe" || request.clientid.is_empty() { return Json(MqttAuthResponse::deny()); } + let Ok(reporter_id) = request.clientid.parse::() else { + return Json(MqttAuthResponse::deny()); + }; + let active = + sqlx::query_scalar::<_, bool>("SELECT EXISTS(SELECT 1 FROM reporters WHERE id = $1)") + .bind(reporter_id) + .fetch_optional(&**pool) + .await + .ok() + .flatten() + .unwrap_or(false); + if !active { + return Json(MqttAuthResponse::deny()); + } match request.access { 1 if can_probe_subscribe(request.clientid, request.topic, pool).await => {