diff --git a/src/control/firehose.rs b/src/control/firehose.rs index 8f5161c..07f74ad 100644 --- a/src/control/firehose.rs +++ b/src/control/firehose.rs @@ -1,9 +1,13 @@ use std::sync::Arc; +#[cfg(test)] +use std::sync::Mutex; use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; use miette::{IntoDiagnostic, Result}; use rand::RngExt; +#[cfg(test)] +use tokio::sync::Notify; use tokio_util::sync::CancellationToken; use tracing::{debug, error, info}; use url::Url; @@ -119,6 +123,8 @@ pub struct FirehoseHandle { pub(super) known_sources: Arc>, /// ids assigned to spawned tasks next_task_id: Arc, + #[cfg(test)] + admission_pause: Arc, Arc)>>>, } impl FirehoseHandle { @@ -129,6 +135,8 @@ impl FirehoseHandle { tasks: Arc::new(scc::HashMap::new()), known_sources: Arc::new(scc::HashMap::new()), next_task_id: Arc::new(AtomicUsize::new(0)), + #[cfg(test)] + admission_pause: Arc::new(Mutex::new(None)), } } @@ -246,6 +254,29 @@ impl FirehoseHandle { ); } + #[cfg(test)] + pub(crate) fn pause_validated_admission_after_ban_check(&self) -> (Arc, Arc) { + let pause = (Arc::new(Notify::new()), Arc::new(Notify::new())); + *self + .admission_pause + .lock() + .expect("admission test pause lock must not be poisoned") = Some(pause.clone()); + pause + } + + #[cfg(test)] + async fn pause_validated_admission_after_ban_check_if_requested(&self) { + let pause = self + .admission_pause + .lock() + .expect("admission test pause lock must not be poisoned") + .clone(); + if let Some((checked, resume)) = pause { + checked.notify_one(); + resume.notified().await; + } + } + /// enable firehose ingestion, no-op if already enabled. pub fn enable(&self) { self.state.firehose_enabled.send_replace(true); @@ -370,6 +401,9 @@ impl FirehoseHandle { if self.state.pds_meta.load().is_banned(host.as_str()) { return Ok(Err(PdsAdmissionRejection::HostBanned)); } + #[cfg(test)] + self.pause_validated_admission_after_ban_check_if_requested() + .await; if self.is_source_running(&url) { return Ok(Ok(PdsAdmission::AlreadyRunning)); } @@ -398,9 +432,19 @@ impl FirehoseHandle { } }; - if let Err(error) = self.persist_validated_source(&source, reservation).await { - self.state.pds_daily_limit.refund(reservation); - return Err(error); + match self + .persist_validated_source(&source, host.as_str(), reservation) + .await + { + Ok(true) => {} + Ok(false) => { + self.state.pds_daily_limit.refund(reservation); + return Ok(Err(PdsAdmissionRejection::HostBanned)); + } + Err(error) => { + self.state.pds_daily_limit.refund(reservation); + return Err(error); + } } let _ = self @@ -417,9 +461,11 @@ impl FirehoseHandle { async fn persist_validated_source( &self, source: &FirehoseSource, + host: &str, reservation: PdsDailyReservation, - ) -> Result<()> { + ) -> Result { let key = keys::firehose_source_key(source.url.as_str()); + let status_key = db::pds_meta::pds_status_key(host); let value = rmp_serde::to_vec(&db::FirehoseSourceMeta { is_pds: true }).map_err(|error| { miette::miette!("failed to serialize firehose source meta: {error}") @@ -427,13 +473,28 @@ impl FirehoseHandle { self.state .db .run(move |db| { + let banned = db + .filter + .get(status_key) + .into_diagnostic()? + .map(|status| { + rmp_serde::from_slice::(status.as_ref()) + .into_diagnostic() + }) + .transpose()? + .is_some_and(|status| status == crate::pds_meta::HostStatus::Banned); + if banned { + return Ok(false); + } + let mut batch = db.inner.batch(); if let Some((day, count)) = reservation.persisted() { db::stage_pds_daily_adds(&mut batch, db, day, count); } batch.insert(&db.crawler, key, value); batch.commit().into_diagnostic()?; - db.persist() + db.persist()?; + Ok(true) }) .await } @@ -725,4 +786,86 @@ mod tests { hydrant.firehose.tasks.remove_async(&url).await; Ok(()) } + + #[tokio::test] + async fn durable_ban_prevents_validated_admission() -> Result<()> { + let dir = tempfile::tempdir().into_diagnostic()?; + let hydrant = Hydrant::new(test_config_with_limit(&dir, 1)).await?; + startable(&hydrant); + let host = "durably-banned.example"; + let url = Url::parse("wss://durably-banned.example/").into_diagnostic()?; + + hydrant + .state + .db + .run(move |db| { + let mut batch = db.inner.batch(); + crate::db::pds_meta::set_status( + &mut batch, + &db.filter, + host, + crate::pds_meta::HostStatus::Banned, + )?; + batch.commit().into_diagnostic()?; + db.persist() + }) + .await?; + assert!( + !hydrant.pds.is_banned(host), + "the durable check is independent of the RCU snapshot" + ); + + assert_eq!( + hydrant.firehose.add_source_validated(url.clone()).await?, + Err(PdsAdmissionRejection::HostBanned) + ); + assert!(!hydrant.firehose.is_source_known(&url)); + assert!(!hydrant.firehose.is_source_running(&url)); + assert!(persisted_sources(&hydrant).await?.is_empty()); + assert_eq!(crate::db::load_pds_daily_adds(&hydrant.state.db)?, None); + Ok(()) + } + + #[tokio::test] + async fn inflight_admission_starts_before_its_ban_is_published() -> Result<()> { + let dir = tempfile::tempdir().into_diagnostic()?; + let hydrant = Hydrant::new(test_config(&dir)).await?; + startable(&hydrant); + let host = "inflight-ban.example"; + let url = Url::parse("wss://inflight-ban.example/").into_diagnostic()?; + let (checked, resume) = hydrant.firehose.pause_validated_admission_after_ban_check(); + + let admission = tokio::spawn({ + let firehose = hydrant.firehose.clone(); + let url = url.clone(); + async move { firehose.add_source_validated(url).await } + }); + checked.notified().await; + + let mut ban = tokio::spawn({ + let pds = hydrant.pds.clone(); + async move { pds.ban(host).await } + }); + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut ban) + .await + .is_err(), + "the ban cannot publish while admission owns its lock" + ); + + resume.notify_one(); + assert_eq!( + admission + .await + .map_err(|error| miette::miette!("admission task failed: {error}"))??, + Ok(PdsAdmission::StartedNew) + ); + assert!(hydrant.firehose.is_source_running(&url)); + + ban.await + .map_err(|error| miette::miette!("ban task failed: {error}"))??; + assert!(hydrant.pds.is_banned(host)); + hydrant.firehose.tasks.remove_async(&url).await; + Ok(()) + } } diff --git a/src/control/pds.rs b/src/control/pds.rs index ca80bbc..9233983 100644 --- a/src/control/pds.rs +++ b/src/control/pds.rs @@ -9,6 +9,7 @@ use tracing::debug; use crate::config::RateTier; use crate::db::keys::pds_account_count_key; use crate::db::pds_meta as db_pds; +use crate::pds_discovery::canonical_pds_host; use crate::pds_meta::{HostDesc, HostStatus, PdsMeta}; use crate::state::AppState; @@ -47,6 +48,15 @@ pub struct PdsControl(pub(super) Arc); impl PdsControl { async fn update(&self, db_op: F, mem_op: G) -> Result<()> + where + F: FnOnce(&mut fjall::OwnedWriteBatch, &fjall::Keyspace) + Send + 'static, + G: FnOnce(&mut PdsMeta), + { + let _admission = self.0.pds_admission.lock().await; + self.update_locked(db_op, mem_op).await + } + + async fn update_locked(&self, db_op: F, mem_op: G) -> Result<()> where F: FnOnce(&mut fjall::OwnedWriteBatch, &fjall::Keyspace) + Send + 'static, G: FnOnce(&mut PdsMeta), @@ -131,6 +141,7 @@ impl PdsControl { ); } + let _admission = self.0.pds_admission.lock().await; let host = host.as_ref().to_string(); let host_clone = host.clone(); let tier_clone = tier.clone(); @@ -148,7 +159,7 @@ impl PdsControl { let current_status = self.0.pds_meta.load().status(&host); let maybe_status = current_status.check_limit_transition(count, new_tier_limit); - self.update( + self.update_locked( move |batch, ks| { db_pds::set_tier(batch, ks, &host_clone, &tier_clone); if let Some(status) = maybe_status { @@ -169,6 +180,7 @@ impl PdsControl { /// remove any explicit tier assignment for `host`, reverting it to the matched rule or default. pub async fn remove_tier(&self, host: impl AsRef) -> Result<()> { + let _admission = self.0.pds_admission.lock().await; let host = host.as_ref().to_string(); let host_clone = host.clone(); @@ -187,7 +199,7 @@ impl PdsControl { "remove_tier: computed status transition" ); - self.update( + self.update_locked( move |batch, ks| { db_pds::remove_tier(batch, ks, &host_clone); if let Some(status) = maybe_status { @@ -208,7 +220,10 @@ impl PdsControl { /// ban `host` pub async fn ban(&self, host: impl AsRef) -> Result<()> { - let host = host.as_ref().to_string(); + let host = canonical_pds_host(host.as_ref()) + .map_err(|error| miette::miette!("invalid PDS host: {error}"))? + .as_str() + .to_string(); let host_clone = host.clone(); self.update( move |batch, ks| { @@ -225,7 +240,10 @@ impl PdsControl { /// unban `host` pub async fn unban(&self, host: impl AsRef) -> Result<()> { - let host = host.as_ref().to_string(); + let host = canonical_pds_host(host.as_ref()) + .map_err(|error| miette::miette!("invalid PDS host: {error}"))? + .as_str() + .to_string(); let host_clone = host.clone(); self.update( move |batch, ks| db_pds::remove_status(batch, ks, &host_clone),