mirror of
https://github.com/LowderPlay/cheburcheck.git
synced 2026-10-01 19:28:08 +03:00
feat: individual tasks (#86)
* feat: hidden nodes and individual tasks * chore: bump version
This commit is contained in:
@@ -26,9 +26,10 @@ pub struct ProbeReporterInfo {
|
||||
pub asn: Option<String>,
|
||||
}
|
||||
|
||||
#[get("/probe/<id>")]
|
||||
#[get("/probe/<id>?<token>")]
|
||||
pub async fn probe_query(
|
||||
id: &str,
|
||||
token: Option<&str>,
|
||||
addr: &ClientRealAddr,
|
||||
pool: &State<PgPool>,
|
||||
mqtt: &State<MqttPublisher>,
|
||||
@@ -67,12 +68,17 @@ pub async fn probe_query(
|
||||
return Err(Status::Forbidden);
|
||||
}
|
||||
|
||||
let (target_probe, eligible_probes) = load_probe_targets(pool, token).await?;
|
||||
let expected_probes = mqtt.online_probe_ids(&eligible_probes).await;
|
||||
let online_probes = expected_probes.len();
|
||||
let expected_probes = expected_probes.into_iter().collect::<HashSet<_>>();
|
||||
|
||||
let mut results = mqtt.subscribe_probe_results(id).await.map_err(|error| {
|
||||
warn!("api: failed to subscribe to probe results for {id}: {error}");
|
||||
publish_error_status(error)
|
||||
})?;
|
||||
|
||||
mqtt.publish_probe_task(id, domain, ip)
|
||||
mqtt.publish_probe_task(id, domain, ip, target_probe.as_deref())
|
||||
.await
|
||||
.map_err(|error| {
|
||||
warn!("api: failed to publish probe task for {id}: {error}");
|
||||
@@ -80,7 +86,6 @@ pub async fn probe_query(
|
||||
})?;
|
||||
|
||||
let timeout = mqtt.task_timeout();
|
||||
let online_probes = mqtt.online_probe_count().await;
|
||||
let pool = pool.inner().clone();
|
||||
let query_id = id;
|
||||
let id = id.to_string();
|
||||
@@ -105,6 +110,9 @@ pub async fn probe_query(
|
||||
result = results.recv() => {
|
||||
match result {
|
||||
Ok(result) => {
|
||||
if !expected_probes.contains(&result.probe_id) {
|
||||
continue;
|
||||
}
|
||||
responded_probes.insert(result.probe_id.clone());
|
||||
let target_traceroute = result.target_traceroute.clone();
|
||||
let reporter_info = match fetch_probe_reporter_info(&result.probe_id, &pool).await {
|
||||
@@ -145,6 +153,44 @@ pub async fn probe_query(
|
||||
})
|
||||
}
|
||||
|
||||
async fn load_probe_targets(
|
||||
pool: &PgPool,
|
||||
token: Option<&str>,
|
||||
) -> Result<(Option<String>, Vec<String>), Status> {
|
||||
if token.is_some_and(str::is_empty) {
|
||||
return Err(Status::BadRequest);
|
||||
}
|
||||
if let Some(token) = token {
|
||||
let probe_id =
|
||||
sqlx::query_scalar::<_, i32>("SELECT id FROM reporters WHERE token = $1 LIMIT 1")
|
||||
.bind(token)
|
||||
.fetch_optional(pool)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
warn!("api: failed to resolve targeted probe: {error}");
|
||||
Status::InternalServerError
|
||||
})?
|
||||
.ok_or(Status::NotFound)?
|
||||
.to_string();
|
||||
|
||||
return Ok((Some(probe_id.clone()), vec![probe_id]));
|
||||
}
|
||||
|
||||
let probe_ids =
|
||||
sqlx::query_scalar::<_, i32>("SELECT id FROM reporters WHERE hidden = FALSE ORDER BY id")
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
warn!("api: failed to load global probe recipients: {error}");
|
||||
Status::InternalServerError
|
||||
})?
|
||||
.into_iter()
|
||||
.map(|id| id.to_string())
|
||||
.collect();
|
||||
|
||||
Ok((None, probe_ids))
|
||||
}
|
||||
|
||||
pub fn build_probe_response(
|
||||
raw: ProbeResultEvent,
|
||||
config: &ProbeConfig,
|
||||
|
||||
+16
-8
@@ -197,13 +197,17 @@ impl MqttPublisher {
|
||||
Duration::from_millis(self.task_timeout_ms)
|
||||
}
|
||||
|
||||
pub async fn online_probe_count(&self) -> usize {
|
||||
self.probe_statuses
|
||||
.read()
|
||||
.await
|
||||
.values()
|
||||
.filter(|status| status.online)
|
||||
.count()
|
||||
pub async fn online_probe_ids(&self, probe_ids: &[String]) -> Vec<String> {
|
||||
let statuses = self.probe_statuses.read().await;
|
||||
probe_ids
|
||||
.iter()
|
||||
.filter(|probe_id| {
|
||||
statuses
|
||||
.get(probe_id.as_str())
|
||||
.is_some_and(|status| status.online)
|
||||
})
|
||||
.cloned()
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub async fn probe_statuses(&self) -> HashMap<String, ProbeStatusSnapshot> {
|
||||
@@ -253,6 +257,7 @@ impl MqttPublisher {
|
||||
query_id: Uuid,
|
||||
domain: Option<&str>,
|
||||
ip: IpAddr,
|
||||
probe_id: Option<&str>,
|
||||
) -> Result<(), PublishError> {
|
||||
let client = self.client.as_ref().ok_or(PublishError::NotConfigured)?;
|
||||
let query_id = query_id.to_string();
|
||||
@@ -265,7 +270,10 @@ impl MqttPublisher {
|
||||
timeout_ms: self.task_timeout_ms,
|
||||
};
|
||||
let payload = serde_json::to_vec(&task).map_err(PublishError::Serialize)?;
|
||||
let topic = format!("probe/tasks/v1/{query_id}");
|
||||
let topic = match probe_id {
|
||||
Some(probe_id) => format!("probe/tasks/v1/{probe_id}/{query_id}"),
|
||||
None => format!("probe/tasks/v1/{query_id}"),
|
||||
};
|
||||
|
||||
client
|
||||
.publish(topic, QoS::AtLeastOnce, false, payload)
|
||||
|
||||
@@ -95,7 +95,10 @@ pub async fn auth(
|
||||
}
|
||||
|
||||
#[post("/acl", data = "<request>")]
|
||||
pub async fn acl(request: Form<MqttAclRequest<'_>>) -> Json<MqttAuthResponse> {
|
||||
pub async fn acl(
|
||||
request: Form<MqttAclRequest<'_>>,
|
||||
pool: &rocket::State<PgPool>,
|
||||
) -> Json<MqttAuthResponse> {
|
||||
let request = request.into_inner();
|
||||
let _ = request.protocol;
|
||||
|
||||
@@ -108,17 +111,40 @@ pub async fn acl(request: Form<MqttAclRequest<'_>>) -> Json<MqttAuthResponse> {
|
||||
}
|
||||
|
||||
match request.access {
|
||||
1 if can_probe_subscribe(request.topic) => Json(MqttAuthResponse::allow()),
|
||||
1 if can_probe_subscribe(request.clientid, request.topic, pool).await => {
|
||||
Json(MqttAuthResponse::allow())
|
||||
}
|
||||
2 if can_probe_publish(request.clientid, request.topic) => Json(MqttAuthResponse::allow()),
|
||||
_ => Json(MqttAuthResponse::deny()),
|
||||
}
|
||||
}
|
||||
|
||||
fn can_probe_subscribe(topic: &str) -> bool {
|
||||
matches!(
|
||||
topic,
|
||||
"probe/config/v1" | "probe/tasks/v1/+" | "probe/tasks/v1/#"
|
||||
)
|
||||
async fn can_probe_subscribe(client_id: &str, topic: &str, pool: &PgPool) -> bool {
|
||||
if topic == "probe/config/v1" || is_own_task_subscription(client_id, topic) {
|
||||
return true;
|
||||
}
|
||||
if !is_global_task_subscription(topic) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let Ok(reporter_id) = client_id.parse::<i32>() else {
|
||||
return false;
|
||||
};
|
||||
sqlx::query_scalar::<_, bool>("SELECT NOT hidden FROM reporters WHERE id = $1")
|
||||
.bind(reporter_id)
|
||||
.fetch_optional(pool)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn is_own_task_subscription(client_id: &str, topic: &str) -> bool {
|
||||
topic == format!("probe/tasks/v1/{client_id}/+")
|
||||
}
|
||||
|
||||
fn is_global_task_subscription(topic: &str) -> bool {
|
||||
topic == "probe/tasks/v1/+"
|
||||
}
|
||||
|
||||
fn can_probe_publish(client_id: &str, topic: &str) -> bool {
|
||||
@@ -141,3 +167,24 @@ fn can_probe_publish(client_id: &str, topic: &str) -> bool {
|
||||
if probe_id == client_id
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn individual_task_subscriptions_are_node_scoped() {
|
||||
assert!(is_global_task_subscription("probe/tasks/v1/+"));
|
||||
assert!(!is_global_task_subscription("probe/tasks/v1/#"));
|
||||
assert!(is_own_task_subscription("42", "probe/tasks/v1/42/+"));
|
||||
assert!(!is_own_task_subscription("42", "probe/tasks/v1/7/+"));
|
||||
assert!(!is_own_task_subscription("42", "probe/tasks/v1/+/+"));
|
||||
assert!(!is_own_task_subscription("42", "probe/tasks/v1/#"));
|
||||
}
|
||||
|
||||
#[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"));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user