diff --git a/examples/statusphere.rs b/examples/statusphere.rs index 60f4dbf..5419e1d 100644 --- a/examples/statusphere.rs +++ b/examples/statusphere.rs @@ -20,9 +20,9 @@ use std::time::Duration; use chrono::DateTime; use futures::StreamExt; +use hydrant::FilterMode; use hydrant::config::Config; use hydrant::control::{EventStream, Hydrant, ReposControl}; -use hydrant::filter::FilterMode; use jacquard_common::types::did::Did; use jacquard_common::types::tid::Tid; use scc::HashMap; diff --git a/src/api/pds.rs b/src/api/pds.rs index ce6baae..808ebc3 100644 --- a/src/api/pds.rs +++ b/src/api/pds.rs @@ -8,7 +8,7 @@ use axum::{ }; use serde::{Deserialize, Serialize}; -use crate::control::{Hydrant, PdsTierAssignment, PdsTierDefinition}; +use crate::control::{Hydrant, PdsTierDefinition}; pub fn router() -> Router { Router::new() @@ -16,22 +16,29 @@ pub fn router() -> Router { .route("/pds/tiers", put(set_tier)) .route("/pds/tiers", delete(remove_tier)) .route("/pds/rate-tiers", get(list_rate_tiers)) + .route("/pds/banned", get(list_banned)) + .route("/pds/banned", put(ban)) + .route("/pds/banned", delete(unban)) } /// combined response: tier assignments + available tier definitions. #[derive(Serialize)] pub struct TiersResponse { - pub assignments: Vec, + pub assignments: HashMap, pub rate_tiers: HashMap, } pub async fn list_tiers(State(hydrant): State) -> Json { Json(TiersResponse { - assignments: hydrant.pds.list_assignments().await, + assignments: hydrant.pds.list_tiers().await, rate_tiers: hydrant.pds.list_rate_tiers(), }) } +pub async fn list_banned(State(hydrant): State) -> Json> { + Json(hydrant.pds.list_banned().await) +} + pub async fn list_rate_tiers( State(hydrant): State, ) -> Json> { @@ -72,3 +79,32 @@ pub async fn remove_tier( .map(|_| StatusCode::OK) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) } + +#[derive(Deserialize)] +pub struct BanBody { + pub host: String, +} + +pub async fn ban( + State(hydrant): State, + Json(body): Json, +) -> Result { + hydrant + .pds + .ban(body.host) + .await + .map(|_| StatusCode::OK) + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) +} + +pub async fn unban( + State(hydrant): State, + Json(body): Json, +) -> Result { + hydrant + .pds + .unban(body.host) + .await + .map(|_| StatusCode::OK) + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) +} diff --git a/src/api/xrpc/get_host_status.rs b/src/api/xrpc/get_host_status.rs index 2f19ee9..e754b39 100644 --- a/src/api/xrpc/get_host_status.rs +++ b/src/api/xrpc/get_host_status.rs @@ -1,5 +1,8 @@ -use jacquard_api::com_atproto::sync::get_host_status::{ - GetHostStatusError, GetHostStatusOutput, GetHostStatusRequest, GetHostStatusResponse, +use jacquard_api::com_atproto::sync::{ + HostStatus, + get_host_status::{ + GetHostStatusError, GetHostStatusOutput, GetHostStatusRequest, GetHostStatusResponse, + }, }; use jacquard_common::CowStr; @@ -23,7 +26,7 @@ pub async fn handle( account_count: Some(host.account_count as i64), hostname: CowStr::Owned(host.name), seq: Some(host.seq), - status: None, + status: host.is_banned.then_some(HostStatus::Banned), extra_data: None, })) } diff --git a/src/api/xrpc/list_hosts.rs b/src/api/xrpc/list_hosts.rs index 64193a9..d83e230 100644 --- a/src/api/xrpc/list_hosts.rs +++ b/src/api/xrpc/list_hosts.rs @@ -1,5 +1,6 @@ -use jacquard_api::com_atproto::sync::list_hosts::{ - Host, ListHostsOutput, ListHostsRequest, ListHostsResponse, +use jacquard_api::com_atproto::sync::{ + HostStatus, + list_hosts::{Host, ListHostsOutput, ListHostsRequest, ListHostsResponse}, }; use jacquard_common::CowStr; @@ -24,7 +25,7 @@ pub async fn handle( .map(|h| Host { hostname: CowStr::Owned(h.name), seq: Some(h.seq), - status: None, + status: h.is_banned.then_some(HostStatus::Banned), account_count: Some(h.account_count as i64), extra_data: None, }) diff --git a/src/control/firehose.rs b/src/control/firehose.rs index 24ab21e..d29421a 100644 --- a/src/control/firehose.rs +++ b/src/control/firehose.rs @@ -5,6 +5,7 @@ use tokio::sync::watch; use tracing::{error, info}; use url::Url; +use crate::config::FirehoseSource; use crate::db::{self, keys}; use crate::ingest::{BufferTx, firehose::FirehoseIngestor}; use crate::state::AppState; @@ -35,46 +36,6 @@ pub struct FirehoseSourceInfo { pub is_pds: bool, } -pub(super) async fn spawn_firehose_ingestor( - relay_url: &Url, - is_pds: bool, - state: &Arc, - shared: &FirehoseShared, - enabled: watch::Receiver, -) -> Result { - use std::sync::atomic::AtomicI64; - - let start = db::get_firehose_cursor(&state.db, relay_url).await?; - // insert into relay_cursors if not already present; existing in-memory cursor takes precedence - let _ = state - .firehose_cursors - .insert_async(relay_url.clone(), AtomicI64::new(start.unwrap_or(0))) - .await; - - info!(relay = %relay_url, is_pds, cursor = ?start, "starting firehose ingestor"); - - let ingestor = FirehoseIngestor::new( - state.clone(), - shared.buffer_tx.clone(), - relay_url.clone(), - is_pds, - state.filter.clone(), - enabled, - shared.verify_signatures, - ) - .await; - - let relay_for_log = relay_url.clone(); - let abort = tokio::spawn(async move { - if let Err(e) = ingestor.run().await { - error!(relay = %relay_for_log, err = %e, "firehose ingestor exited with error"); - } - }) - .abort_handle(); - - Ok(FirehoseIngestorHandle { abort, is_pds }) -} - /// runtime control over the firehose ingestor component. #[derive(Clone)] pub struct FirehoseHandle { @@ -97,6 +58,59 @@ impl FirehoseHandle { } } + pub(super) async fn spawn_firehose_ingestor( + &self, + source: &FirehoseSource, + shared: &FirehoseShared, + ) -> Result<()> { + use std::sync::atomic::AtomicI64; + let state = &self.state; + + let start = db::get_firehose_cursor(&state.db, &source.url).await?; + // insert into relay_cursors if not already present; existing in-memory cursor takes precedence + let _ = state + .firehose_cursors + .insert_async(source.url.clone(), AtomicI64::new(start.unwrap_or(0))) + .await; + + info!(relay = %source.url, source.is_pds, cursor = ?start, "starting firehose ingestor"); + + let enabled = state.firehose_enabled.subscribe(); + let ingestor = FirehoseIngestor::new( + state.clone(), + shared.buffer_tx.clone(), + source.url.clone(), + source.is_pds, + state.filter.clone(), + enabled, + shared.verify_signatures, + ) + .await; + + let abort = tokio::spawn({ + let relay_url = source.url.clone(); + let tasks = self.tasks.clone(); + async move { + if let Err(e) = ingestor.run().await { + error!(relay = %relay_url, err = %e, "firehose ingestor exited with error"); + } else { + // remove from tasks since we shutdown + tasks.remove_async(&relay_url).await; + info!(relay = %relay_url, "firehose shut down!"); + } + } + }) + .abort_handle(); + + let handle = FirehoseIngestorHandle { + abort, + is_pds: source.is_pds, + }; + self.tasks.upsert_async(source.url.clone(), handle).await; + + Ok(()) + } + /// enable firehose ingestion, no-op if already enabled. pub fn enable(&self) { self.state.firehose_enabled.send_replace(true); @@ -156,9 +170,8 @@ impl FirehoseHandle { let _ = self.persisted.insert_async(url.clone()).await; - let enabled_rx = self.state.firehose_enabled.subscribe(); - let handle = spawn_firehose_ingestor(&url, is_pds, &self.state, shared, enabled_rx).await?; - self.tasks.upsert_async(url, handle).await; + self.spawn_firehose_ingestor(&FirehoseSource { url, is_pds }, shared) + .await?; Ok(()) } diff --git a/src/control/mod.rs b/src/control/mod.rs index 29102eb..e56c777 100644 --- a/src/control/mod.rs +++ b/src/control/mod.rs @@ -46,10 +46,11 @@ use crate::db::{self, filter as db_filter, keys, load_persisted_firehose_sources use crate::filter::FilterMode; #[cfg(feature = "indexer")] use crate::ingest::indexer::FirehoseWorker; +use crate::pds_meta::{PdsMeta, PdsMetaHandle}; use crate::state::AppState; use crate::types::MarshallableEvt; -use firehose::{FirehoseShared, spawn_firehose_ingestor}; +use firehose::FirehoseShared; #[cfg(feature = "indexer")] use stream::event_stream_thread; #[cfg(feature = "relay")] @@ -63,6 +64,8 @@ pub struct Host { pub seq: i64, /// the amount of accounts hydrant has seen from this host. pub account_count: u64, + /// whether this host is banned or not. + pub is_banned: bool, } /// an event emitted by the hydrant event stream. @@ -190,6 +193,11 @@ impl Hydrant { let state = Arc::new(state); Ok(Self { + firehose: FirehoseHandle::new(state.clone()), + filter: FilterControl(state.clone()), + pds: pds::PdsControl(state.clone()), + repos: ReposControl(state.clone()), + db: DbControl(state.clone()), #[cfg(feature = "indexer")] crawler: crawler::CrawlerHandle { state: state.clone(), @@ -197,13 +205,8 @@ impl Hydrant { tasks: Arc::new(scc::HashMap::new()), persisted: Arc::new(scc::HashSet::new()), }, - firehose: FirehoseHandle::new(state.clone()), #[cfg(feature = "indexer")] backfill: BackfillHandle::new(state.clone()), - filter: FilterControl(state.clone()), - pds: pds::PdsControl(state.clone()), - repos: ReposControl(state.clone()), - db: DbControl(state.clone()), #[cfg(feature = "backlinks")] backlinks: crate::backlinks::BacklinksControl(state.clone()), state, @@ -404,19 +407,9 @@ impl Hydrant { "starting firehose ingestor(s)" ); for source in &relay_hosts { - let enabled_rx = state.firehose_enabled.subscribe(); - let handle = spawn_firehose_ingestor( - &source.url, - source.is_pds, - &state, - fire_shared, - enabled_rx, - ) - .await?; - let _ = firehose - .tasks - .insert_async(source.url.clone(), handle) - .await; + firehose + .spawn_firehose_ingestor(source, fire_shared) + .await?; } } @@ -432,19 +425,9 @@ impl Hydrant { if firehose.tasks.contains_async(&source.url).await { continue; } - let enabled_rx = state.firehose_enabled.subscribe(); - let handle = spawn_firehose_ingestor( - &source.url, - source.is_pds, - &state, - fire_shared, - enabled_rx, - ) - .await?; - let _ = firehose - .tasks - .insert_async(source.url.clone(), handle) - .await; + firehose + .spawn_firehose_ingestor(source, fire_shared) + .await?; } // 10c. seed firehose PDS sources from listHosts on configured seed URLs @@ -796,11 +779,13 @@ impl Hydrant { let account_count = state .db .get_count_sync(&keys::pds_account_count_key(&hostname)); + let is_banned = state.pds_meta.load().is_banned(&hostname); Ok(Some(Host { name: hostname.into(), seq, account_count, + is_banned, })) }) .await @@ -850,10 +835,12 @@ impl Hydrant { let account_count = state .db .get_count_sync(&keys::pds_account_count_key(hostname)); + let is_banned = state.pds_meta.load().is_banned(&hostname); hosts.push(Host { name: hostname.into(), seq, account_count, + is_banned, }); } diff --git a/src/control/pds.rs b/src/control/pds.rs index 68ff615..9b7558f 100644 --- a/src/control/pds.rs +++ b/src/control/pds.rs @@ -6,7 +6,8 @@ use serde::Serialize; use smol_str::SmolStr; use crate::config::RateTier; -use crate::db::pds_tiers as db_pds; +use crate::db::pds_meta as db_pds; +use crate::pds_meta::PdsMeta; use crate::state::AppState; /// a single PDS-to-tier assignment. @@ -41,18 +42,59 @@ impl From for PdsTierDefinition { 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 state = self.0.clone(); + tokio::task::spawn_blocking(move || { + let mut batch = state.db.inner.batch(); + db_op(&mut batch, &state.db.filter); + batch.commit().into_diagnostic()?; + state.db.persist() + }) + .await + .into_diagnostic()??; + + let mut snapshot = (**self.0.pds_meta.load()).clone(); + mem_op(&mut snapshot); + self.0.pds_meta.store(Arc::new(snapshot)); + + Ok(()) + } + /// list all current per-PDS tier assignments. - pub async fn list_assignments(&self) -> Vec { - let snapshot = self.0.pds_tiers.load(); + pub async fn list_tiers(&self) -> HashMap { + let snapshot = self.0.pds_meta.load(); snapshot + .tiers .iter() - .map(|(host, tier)| PdsTierAssignment { - host: host.clone(), - tier: tier.to_string(), - }) + .map(|(host, tier)| (host.clone(), tier.to_string())) .collect() } + /// returns the assigned tier for `host`, or "default" if none is assigned. + pub fn get_tier(&self, host: impl AsRef) -> String { + let snapshot = self.0.pds_meta.load(); + snapshot + .tiers + .get(host.as_ref()) + .map(|t| t.to_string()) + .unwrap_or_else(|| "default".to_string()) + } + + /// returns true if `host` is currently banned. + pub fn is_banned(&self, host: impl AsRef) -> bool { + self.0.pds_meta.load().is_banned(host.as_ref()) + } + + /// list all currently banned PDS hosts. + pub async fn list_banned(&self) -> Vec { + let snapshot = self.0.pds_meta.load(); + snapshot.banned.iter().cloned().collect() + } + /// list all configured rate tier definitions. pub fn list_rate_tiers(&self) -> HashMap { self.0 @@ -64,7 +106,7 @@ impl PdsControl { /// assign `host` to `tier`, persisting the change to the database. /// returns an error if `tier` is not a known tier name. - pub async fn set_tier(&self, host: String, tier: String) -> Result<()> { + pub async fn set_tier(&self, host: impl AsRef, tier: String) -> Result<()> { if !self.0.rate_tiers.contains_key(&tier) { miette::bail!( "unknown tier '{tier}'; known tiers: {:?}", @@ -72,42 +114,54 @@ impl PdsControl { ); } - let state = self.0.clone(); + let host = host.as_ref().to_string(); let host_clone = host.clone(); let tier_clone = tier.clone(); - tokio::task::spawn_blocking(move || { - let mut batch = state.db.inner.batch(); - db_pds::set(&mut batch, &state.db.filter, &host_clone, &tier_clone); - batch.commit().into_diagnostic()?; - state.db.persist() - }) + self.update( + move |batch, ks| db_pds::set_tier(batch, ks, &host_clone, &tier_clone), + move |meta| { + meta.tiers.insert(host, SmolStr::new(&tier)); + }, + ) .await - .into_diagnostic()??; - - let mut snapshot = (**self.0.pds_tiers.load()).clone(); - snapshot.insert(host, SmolStr::new(&tier)); - self.0.pds_tiers.store(Arc::new(snapshot)); - - Ok(()) } /// remove any explicit tier assignment for `host`, reverting it to the default tier. - pub async fn remove_tier(&self, host: String) -> Result<()> { - let state = self.0.clone(); + pub async fn remove_tier(&self, host: impl AsRef) -> Result<()> { + let host = host.as_ref().to_string(); let host_clone = host.clone(); - tokio::task::spawn_blocking(move || { - let mut batch = state.db.inner.batch(); - db_pds::remove(&mut batch, &state.db.filter, &host_clone); - batch.commit().into_diagnostic()?; - state.db.persist() - }) + self.update( + move |batch, ks| db_pds::remove_tier(batch, ks, &host_clone), + move |meta| { + meta.tiers.remove(&host); + }, + ) .await - .into_diagnostic()??; + } - let mut snapshot = (**self.0.pds_tiers.load()).clone(); - snapshot.remove(&host); - self.0.pds_tiers.store(Arc::new(snapshot)); + /// ban `host`, persisting the change to the database. + pub async fn ban(&self, host: impl AsRef) -> Result<()> { + let host = host.as_ref().to_string(); + let host_clone = host.clone(); + self.update( + move |batch, ks| db_pds::set_banned(batch, ks, &host_clone), + move |meta| { + meta.banned.insert(host); + }, + ) + .await + } - Ok(()) + /// unban `host`, removing it from the database. + pub async fn unban(&self, host: impl AsRef) -> Result<()> { + let host = host.as_ref().to_string(); + let host_clone = host.clone(); + self.update( + move |batch, ks| db_pds::remove_banned(batch, ks, &host_clone), + move |meta| { + meta.banned.remove(&host); + }, + ) + .await } } diff --git a/src/db/mod.rs b/src/db/mod.rs index 51441d9..76b6976 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -27,7 +27,7 @@ pub mod ephemeral; pub mod filter; pub mod keys; pub mod migration; -pub mod pds_tiers; +pub mod pds_meta; pub mod types; use tokio::sync::broadcast; @@ -389,7 +389,7 @@ impl Db { opts() // only iterators are used here .expect_point_read_hits(true) - .max_memtable_size(mb(16)) + .max_memtable_size(mb(8)) // did -> failed state, not very compressable .data_block_size_policy(BlockSizePolicy::all(kb(2))) .data_block_compression_policy(CompressionPolicy::disabled()) diff --git a/src/db/pds_meta.rs b/src/db/pds_meta.rs new file mode 100644 index 0000000..7e9489c --- /dev/null +++ b/src/db/pds_meta.rs @@ -0,0 +1,62 @@ +use fjall::{Keyspace, OwnedWriteBatch}; +use miette::{IntoDiagnostic, Result}; +use smol_str::SmolStr; + +pub const PDS_TIER_PREFIX: &[u8] = b"pt|"; + +// `pt|{host}` -> tier name +pub fn pds_tier_key(host: &str) -> Vec { + let mut key = Vec::with_capacity(PDS_TIER_PREFIX.len() + host.len()); + key.extend_from_slice(PDS_TIER_PREFIX); + key.extend_from_slice(host.as_bytes()); + key +} + +/// load all PDS tier assignments from the filter keyspace +pub fn load_tiers(ks: &Keyspace) -> Result> { + let mut out = Vec::new(); + for guard in ks.prefix(PDS_TIER_PREFIX) { + let (k, v) = guard.into_inner().into_diagnostic()?; + let host = std::str::from_utf8(&k[PDS_TIER_PREFIX.len()..]).into_diagnostic()?; + let tier = std::str::from_utf8(&v).into_diagnostic()?; + out.push((SmolStr::new(host), SmolStr::new(tier))); + } + Ok(out) +} + +pub fn set_tier(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str, tier: &str) { + batch.insert(ks, pds_tier_key(host), tier.as_bytes()); +} + +pub fn remove_tier(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str) { + batch.remove(ks, pds_tier_key(host)); +} + +pub const PDS_BANNED_PREFIX: &[u8] = b"pb|"; + +// `pb|{host}` -> empty value +pub fn pds_banned_key(host: &str) -> Vec { + let mut key = Vec::with_capacity(PDS_BANNED_PREFIX.len() + host.len()); + key.extend_from_slice(PDS_BANNED_PREFIX); + key.extend_from_slice(host.as_bytes()); + key +} + +/// load all banned PDS hosts from the filter keyspace +pub fn load_banned(ks: &Keyspace) -> Result> { + let mut out = Vec::new(); + for guard in ks.prefix(PDS_BANNED_PREFIX) { + let (k, _) = guard.into_inner().into_diagnostic()?; + let host = std::str::from_utf8(&k[PDS_BANNED_PREFIX.len()..]).into_diagnostic()?; + out.push(SmolStr::new(host)); + } + Ok(out) +} + +pub fn set_banned(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str) { + batch.insert(ks, pds_banned_key(host), &[]); +} + +pub fn remove_banned(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str) { + batch.remove(ks, pds_banned_key(host)); +} diff --git a/src/db/pds_tiers.rs b/src/db/pds_tiers.rs deleted file mode 100644 index 5ae44d6..0000000 --- a/src/db/pds_tiers.rs +++ /dev/null @@ -1,33 +0,0 @@ -use fjall::{Keyspace, OwnedWriteBatch}; -use miette::{IntoDiagnostic, Result}; -use smol_str::SmolStr; - -pub const PDS_TIER_PREFIX: &[u8] = b"pt|"; - -// `pt|{host}` -> tier name -pub fn pds_tier_key(host: &str) -> Vec { - let mut key = Vec::with_capacity(PDS_TIER_PREFIX.len() + host.len()); - key.extend_from_slice(PDS_TIER_PREFIX); - key.extend_from_slice(host.as_bytes()); - key -} - -/// load all PDS tier assignments from the filter keyspace -pub fn load(ks: &Keyspace) -> Result> { - let mut out = Vec::new(); - for guard in ks.prefix(PDS_TIER_PREFIX) { - let (k, v) = guard.into_inner().into_diagnostic()?; - let host = std::str::from_utf8(&k[PDS_TIER_PREFIX.len()..]).into_diagnostic()?; - let tier = std::str::from_utf8(&v).into_diagnostic()?; - out.push((SmolStr::new(host), SmolStr::new(tier))); - } - Ok(out) -} - -pub fn set(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str, tier: &str) { - batch.insert(ks, pds_tier_key(host), tier.as_bytes()); -} - -pub fn remove(batch: &mut OwnedWriteBatch, ks: &Keyspace, host: &str) { - batch.remove(ks, pds_tier_key(host)); -} diff --git a/src/ingest/firehose.rs b/src/ingest/firehose.rs index 47e8bc1..dc2db81 100644 --- a/src/ingest/firehose.rs +++ b/src/ingest/firehose.rs @@ -105,9 +105,13 @@ impl FirehoseIngestor { // this is not for connection throttling (thats handled by ThrottleHandle) // its for stream errors (cbor decode etc) let mut backoff = Duration::from_secs(0); - const MAX_BACKOFF: Duration = Duration::from_secs(60 * 60); // 1 ohur + const MAX_BACKOFF: Duration = Duration::from_secs(60 * 60); // 1 hour loop { + if self.state.pds_meta.load().is_banned(host) { + break Ok(()); + } + self.enabled.wait_enabled("firehose").await; tokio::time::sleep(backoff).await; @@ -169,8 +173,15 @@ impl FirehoseIngestor { match decode_frame(&bytes) { Ok(msg) => { if self.is_pds { + let tier = { + let meta = self.state.pds_meta.load(); + let banned = meta.is_banned(host); + if banned { + break Ok(()); + } + meta.tier_for(host, &self.state.rate_tiers) + }; let accounts = self.state.db.get_count(&count_key).await; - let tier = self.state.pds_tier_for(&host); tokio::select! { _ = self.throttle.wait_for_allow(accounts, &tier) => {} _ = self.enabled.changed() => { diff --git a/src/lib.rs b/src/lib.rs index 572b3bb..0680814 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,8 @@ pub mod config; +/// hydrant main api, includes the Hydrant type for programmatic control. pub mod control; -pub mod filter; +pub(crate) mod filter; +pub(crate) mod pds_meta; pub mod types; #[cfg(all(feature = "relay", feature = "indexer"))] @@ -23,3 +25,5 @@ pub(crate) mod patch; pub(crate) mod resolver; pub(crate) mod state; pub(crate) mod util; + +pub use filter::FilterMode; diff --git a/src/pds_meta.rs b/src/pds_meta.rs new file mode 100644 index 0000000..44281b1 --- /dev/null +++ b/src/pds_meta.rs @@ -0,0 +1,34 @@ +use crate::config::RateTier; +use arc_swap::ArcSwap; +use smol_str::SmolStr; +use std::collections::{HashMap, HashSet}; +use std::sync::Arc; + +#[derive(Default, Clone)] +pub(crate) struct PdsMeta { + pub tiers: HashMap, + pub banned: HashSet, +} + +impl PdsMeta { + pub fn tier_for(&self, host: &str, rate_tiers: &HashMap) -> RateTier { + let default = rate_tiers + .get("default") + .copied() + .unwrap_or_else(RateTier::default_tier); + self.tiers + .get(host) + .and_then(|name| rate_tiers.get(name.as_str()).copied()) + .unwrap_or(default) + } + + pub fn is_banned(&self, host: &str) -> bool { + self.banned.contains(host) + } +} + +pub(crate) type PdsMetaHandle = Arc>; + +pub(crate) fn new_handle(meta: PdsMeta) -> PdsMetaHandle { + Arc::new(ArcSwap::new(Arc::new(meta))) +} diff --git a/src/state.rs b/src/state.rs index f602a32..6f56dfb 100644 --- a/src/state.rs +++ b/src/state.rs @@ -1,9 +1,8 @@ -use std::collections::HashMap; -use std::sync::Arc; +use std::collections::{HashMap, HashSet}; +use std::future::Future; use std::sync::atomic::AtomicI64; use std::time::Duration; -use arc_swap::ArcSwap; use miette::Result; use smol_str::SmolStr; #[cfg(feature = "indexer")] @@ -14,19 +13,17 @@ use url::Url; use crate::{ config::{Config, RateTier}, db::Db, - filter::{FilterHandle, new_handle}, + filter::{FilterHandle, new_handle as new_filter_handle}, + pds_meta::{PdsMeta, PdsMetaHandle, new_handle as new_pds_handle}, resolver::Resolver, util::throttle::Throttler, }; -/// pds hostname -> tier name. updated atomically via ArcSwap. -pub(crate) type PdsTierHandle = Arc>>; - pub struct AppState { pub db: Db, pub resolver: Resolver, pub(crate) filter: FilterHandle, - pub(crate) pds_tiers: PdsTierHandle, + pub(crate) pds_meta: PdsMetaHandle, pub(crate) rate_tiers: HashMap, pub firehose_cursors: scc::HashIndex, #[cfg(feature = "indexer")] @@ -56,22 +53,29 @@ impl AppState { } }; - let filter = new_handle(filter_config); + let filter = new_filter_handle(filter_config); // load persisted per-PDS tier assignments from the filter keyspace. // trusted_hosts from config are merged in as defaults (not persisted here; they seed // only if the host has no existing assignment in the DB). - let mut tier_map: HashMap = crate::db::pds_tiers::load(&db.filter) + let mut tiers: HashMap = crate::db::pds_meta::load_tiers(&db.filter) .unwrap_or_default() .into_iter() .map(|(host, tier)| (host.to_string(), tier)) .collect(); for host in &config.trusted_hosts { - tier_map + tiers .entry(host.clone()) .or_insert_with(|| SmolStr::new("trusted")); } - let pds_tiers = Arc::new(ArcSwap::new(Arc::new(tier_map))); + + let banned: HashSet = crate::db::pds_meta::load_banned(&db.filter) + .unwrap_or_default() + .into_iter() + .map(|host| host.to_string()) + .collect(); + + let pds_meta = new_pds_handle(PdsMeta { tiers, banned }); let relay_cursors = scc::HashIndex::new(); @@ -85,7 +89,7 @@ impl AppState { db, resolver, filter, - pds_tiers, + pds_meta, rate_tiers: config.rate_tiers.clone(), firehose_cursors: relay_cursors, #[cfg(feature = "indexer")] @@ -105,21 +109,6 @@ impl AppState { self.backfill_notify.notify_one(); } - /// returns the rate tier for the given PDS hostname. - /// falls back to the "default" tier if no assignment exists or the assigned tier is unknown. - pub fn pds_tier_for(&self, host: &str) -> RateTier { - let default = self - .rate_tiers - .get("default") - .copied() - .unwrap_or_else(RateTier::default_tier); - let snapshot = self.pds_tiers.load(); - snapshot - .get(host) - .and_then(|name| self.rate_tiers.get(name.as_str()).copied()) - .unwrap_or(default) - } - /// pauses the crawler, firehose, and backfill worker, runs `f`, then restores their prior state. /// the restore always happens, even if `f` returns an error. pub async fn with_ingestion_paused(&self, f: F) -> T diff --git a/tests/api.nu b/tests/api.nu index 5d4a1f6..07865d9 100644 --- a/tests/api.nu +++ b/tests/api.nu @@ -256,8 +256,8 @@ def test-pds-tiers [url: string, pid: int] { # initial state: no assignments, built-in rate tiers present print " GET /pds/tiers (expect empty assignments, built-in rate_tiers)..." let initial = (http get $"($url)/pds/tiers") - if ($initial.assignments | length) != 0 { - fail $"expected empty assignments, got ($initial.assignments | length)" $pid + if ($initial.assignments | columns | length) != 0 { + fail $"expected empty assignments, got ($initial.assignments | columns | length)" $pid } if not ("default" in $initial.rate_tiers) { fail "expected 'default' tier in rate_tiers" $pid @@ -291,17 +291,16 @@ def test-pds-tiers [url: string, pid: int] { tier: "trusted" } | assert-status 200 "PUT /pds/tiers" $pid let after_assign = (http get $"($url)/pds/tiers") - if ($after_assign.assignments | length) != 1 { - fail $"expected 1 assignment, got ($after_assign.assignments | length)" $pid + if ($after_assign.assignments | columns | length) != 1 { + fail $"expected 1 assignment, got ($after_assign.assignments | columns | length)" $pid } - let a = ($after_assign.assignments | first) - if $a.host != "pds.example.com" { - fail $"expected host=pds.example.com, got ($a.host)" $pid + if not ("pds.example.com" in $after_assign.assignments) { + fail $"expected host=pds.example.com to be assigned" $pid } - if $a.tier != "trusted" { - fail $"expected tier=trusted, got ($a.tier)" $pid + if ($after_assign.assignments | get "pds.example.com") != "trusted" { + fail $"expected tier=trusted" $pid } - print $" ok: assignment created host=($a.host), tier=($a.tier)" + print $" ok: assignment created host=pds.example.com, tier=trusted" # re-assigning the same host to a different tier updates without creating a duplicate print " PUT /pds/tiers (re-assign to default)..." @@ -310,11 +309,11 @@ def test-pds-tiers [url: string, pid: int] { tier: "default" } | assert-status 200 "PUT /pds/tiers re-assign" $pid let after_reassign = (http get $"($url)/pds/tiers") - if ($after_reassign.assignments | length) != 1 { - fail $"expected 1 assignment after re-assign, got ($after_reassign.assignments | length)" $pid + if ($after_reassign.assignments | columns | length) != 1 { + fail $"expected 1 assignment after re-assign, got ($after_reassign.assignments | columns | length)" $pid } - if ($after_reassign.assignments | first).tier != "default" { - fail $"expected tier=default after re-assign, got (($after_reassign.assignments | first).tier)" $pid + if ($after_reassign.assignments | get "pds.example.com") != "default" { + fail $"expected tier=default after re-assign" $pid } print " ok: re-assign updates tier without creating a duplicate" @@ -325,10 +324,10 @@ def test-pds-tiers [url: string, pid: int] { tier: "nonexistent" } | assert-status 400 "PUT /pds/tiers unknown tier" $pid let after_bad = (http get $"($url)/pds/tiers") - if ($after_bad.assignments | length) != 1 { + if ($after_bad.assignments | columns | length) != 1 { fail "expected assignment count unchanged after rejected request" $pid } - if ($after_bad.assignments | first).tier != "default" { + if ($after_bad.assignments | get "pds.example.com") != "default" { fail "expected tier unchanged after rejected request" $pid } print " ok: unknown tier name rejected with 400, existing assignment unchanged" @@ -340,8 +339,8 @@ def test-pds-tiers [url: string, pid: int] { tier: "trusted" } | assert-status 200 "PUT /pds/tiers second host" $pid let after_second = (http get $"($url)/pds/tiers") - if ($after_second.assignments | length) != 2 { - fail $"expected 2 assignments, got ($after_second.assignments | length)" $pid + if ($after_second.assignments | columns | length) != 2 { + fail $"expected 2 assignments, got ($after_second.assignments | columns | length)" $pid } print " ok: two distinct hosts listed independently" @@ -351,10 +350,10 @@ def test-pds-tiers [url: string, pid: int] { host: "pds.example.com" } | assert-status 200 "DELETE /pds/tiers" $pid let after_del = (http get $"($url)/pds/tiers") - if ($after_del.assignments | length) != 1 { - fail $"expected 1 assignment after delete, got ($after_del.assignments | length)" $pid + if ($after_del.assignments | columns | length) != 1 { + fail $"expected 1 assignment after delete, got ($after_del.assignments | columns | length)" $pid } - if ($after_del.assignments | first).host != "other.example.com" { + if not ("other.example.com" in $after_del.assignments) { fail "expected only other.example.com to remain after delete" $pid } print " ok: correct host removed, other assignment intact" @@ -370,7 +369,7 @@ def test-pds-tiers [url: string, pid: int] { host: "pds.example.com" } | assert-status 200 "DELETE /pds/tiers non-existent" $pid let after_idempotent = (http get $"($url)/pds/tiers") - if ($after_idempotent.assignments | length) != 0 { + if ($after_idempotent.assignments | columns | length) != 0 { fail "expected empty assignments after cleanup" $pid } print " ok: delete of non-existent host is idempotent" @@ -398,7 +397,7 @@ def test-pds-tier-persistence [binary: string, db_path: string, port: int] { } let before = (http get $"($url)/pds/tiers") - if ($before.assignments | length) != 1 { + if ($before.assignments | columns | length) != 1 { fail "assignment was not created" $instance.pid } @@ -415,15 +414,14 @@ def test-pds-tier-persistence [binary: string, db_path: string, port: int] { print " checking assignment survived restart..." let after = (http get $"($url)/pds/tiers") - if ($after.assignments | length) != 1 { - fail $"expected 1 assignment after restart, got ($after.assignments | length)" $instance2.pid + if ($after.assignments | columns | length) != 1 { + fail $"expected 1 assignment after restart, got ($after.assignments | columns | length)" $instance2.pid } - let a = ($after.assignments | first) - if $a.host != "persist.example.com" { - fail $"expected host=persist.example.com after restart, got ($a.host)" $instance2.pid + if not ("persist.example.com" in $after.assignments) { + fail $"expected host=persist.example.com after restart" $instance2.pid } - if $a.tier != "trusted" { - fail $"expected tier=trusted after restart, got ($a.tier)" $instance2.pid + if ($after.assignments | get "persist.example.com") != "trusted" { + fail $"expected tier=trusted after restart" $instance2.pid } print " ok: tier assignment persisted across restart" @@ -455,12 +453,11 @@ def test-pds-trusted-hosts [binary: string, db_path: string, port: int] { let assignments = $tiers.assignments for host in [$host_a, $host_b] { - let match = ($assignments | where host == $host) - if ($match | length) != 1 { - fail $"expected assignment for ($host) from HYDRANT_TRUSTED_HOSTS, got ($assignments)" $instance.pid + if not ($host in $assignments) { + fail $"expected assignment for ($host) from HYDRANT_TRUSTED_HOSTS" $instance.pid } - if ($match | first).tier != "trusted" { - fail $"expected tier=trusted for ($host), got (($match | first).tier)" $instance.pid + if ($assignments | get $host) != "trusted" { + fail $"expected tier=trusted for ($host)" $instance.pid } } print $" ok: ($host_a) and ($host_b) pre-assigned to trusted tier" @@ -512,12 +509,11 @@ def test-pds-custom-rate-tier [binary: string, db_path: string, port: int] { tier: "custom" } | assert-status 200 "PUT /pds/tiers custom tier" $instance.pid let after = (http get $"($url)/pds/tiers") - let match = ($after.assignments | where host == "custom.example.com") - if ($match | length) != 1 { + if not ("custom.example.com" in $after.assignments) { fail "expected assignment for custom.example.com" $instance.pid } - if ($match | first).tier != "custom" { - fail $"expected tier=custom, got (($match | first).tier)" $instance.pid + if ($after.assignments | get "custom.example.com") != "custom" { + fail $"expected tier=custom" $instance.pid } print " ok: host assigned to custom tier successfully" @@ -525,6 +521,92 @@ def test-pds-custom-rate-tier [binary: string, db_path: string, port: int] { print "custom rate tier test passed!" } +def test-pds-banned [url: string, pid: int] { + print "=== test: pds ban management ===" + + print " GET /pds/banned (expect empty)..." + let initial = (http get $"($url)/pds/banned") + if ($initial | length) != 0 { + fail $"expected empty banned list, got ($initial | length)" $pid + } + print " ok: starts empty" + + print " PUT /pds/banned (ban host)..." + http put -f -e -t application/json $"($url)/pds/banned" { + host: "bad.example.com" + } | assert-status 200 "PUT /pds/banned" $pid + + let after_ban = (http get $"($url)/pds/banned") + if ($after_ban | length) != 1 { + fail $"expected 1 banned host, got ($after_ban | length)" $pid + } + if ($after_ban | first) != "bad.example.com" { + fail $"expected bad.example.com, got ($after_ban | first)" $pid + } + print " ok: host banned" + + print " DELETE /pds/banned (unban host)..." + http delete -f -e -t application/json $"($url)/pds/banned" --data { + host: "bad.example.com" + } | assert-status 200 "DELETE /pds/banned" $pid + + let after_unban = (http get $"($url)/pds/banned") + if ($after_unban | length) != 0 { + fail "expected empty banned list after unban" $pid + } + print " ok: host unbanned" + + print "pds ban management tests passed!" +} + +# verify that banned hosts are written to the database and survive a restart. +def test-pds-banned-persistence [binary: string, db_path: string, port: int] { + print "=== test: pds banned assignments persist across restart ===" + + let url = $"http://localhost:($port)" + + let instance = (with-env { HYDRANT_CRAWLER_URLS: "", HYDRANT_RELAY_HOSTS: "" } { + start-hydrant $binary $db_path $port + }) + if not (wait-for-api $url) { + fail "hydrant did not start" + } + + print " banning host..." + http put -t application/json $"($url)/pds/banned" { + host: "persist-ban.example.com" + } + + let before = (http get $"($url)/pds/banned") + if ($before | length) != 1 { + fail "host was not banned" $instance.pid + } + + print " restarting hydrant..." + kill $instance.pid + sleep 2sec + + let instance2 = (with-env { HYDRANT_CRAWLER_URLS: "", HYDRANT_RELAY_HOSTS: "" } { + start-hydrant $binary $db_path $port + }) + if not (wait-for-api $url) { + fail "hydrant did not restart" $instance2.pid + } + + print " checking ban survived restart..." + let after = (http get $"($url)/pds/banned") + if ($after | length) != 1 { + fail $"expected 1 banned host after restart, got ($after | length)" $instance2.pid + } + if ($after | first) != "persist-ban.example.com" { + fail $"expected persist-ban.example.com after restart, got ($after | first)" $instance2.pid + } + print " ok: banned host persisted across restart" + + kill $instance2.pid + print "pds ban persistence test passed!" +} + def main [] { let port = resolve-test-port 3007 let url = $"http://localhost:($port)" @@ -544,6 +626,7 @@ def main [] { test-crawler-sources $url $instance.pid test-firehose-sources $url $instance.pid test-pds-tiers $url $instance.pid + test-pds-banned $url $instance.pid kill $instance.pid sleep 2sec @@ -576,6 +659,12 @@ def main [] { print $"db: ($db_pds_custom)" test-pds-custom-rate-tier $binary $db_pds_custom $port + sleep 1sec + + let db_pds_banned = (mktemp -d -t hydrant_api.XXXXXX) + print $"db: ($db_pds_banned)" + test-pds-banned-persistence $binary $db_pds_banned $port + print "" print "all api tests passed!" }