lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

source_ingress.rs (15260B)


      1 //! Bounded fetch admission at the shared socket's pre-Nostr decoding boundary.
      2 
      3 use crate::source::budget::FetchBudget;
      4 use futures::task::AtomicWaker;
      5 use std::{
      6     collections::{BTreeMap, BTreeSet},
      7     sync::{Arc, Mutex},
      8     task::Context,
      9 };
     10 
     11 const MAX_ACTIVE_FETCHES: usize = 64;
     12 
     13 #[derive(Debug, Default)]
     14 struct RelayState {
     15     generation: u64,
     16     invalidated: u64,
     17     waker: AtomicWaker,
     18     write_waker: AtomicWaker,
     19 }
     20 
     21 #[derive(Debug)]
     22 struct Registration {
     23     targets: BTreeSet<String>,
     24     budget: Arc<FetchBudget>,
     25 }
     26 
     27 #[derive(Debug, Default)]
     28 struct State {
     29     sequence: u64,
     30     relays: BTreeMap<String, RelayState>,
     31     active: BTreeMap<u64, Registration>,
     32 }
     33 
     34 #[derive(Clone, Debug, Default)]
     35 pub(crate) struct IngressRegistry(Arc<Mutex<State>>);
     36 
     37 impl IngressRegistry {
     38     pub(crate) fn new(targets: impl Iterator<Item = String>) -> Self {
     39         Self(Arc::new(Mutex::new(State {
     40             relays: targets
     41                 .map(|target| (target, RelayState::default()))
     42                 .collect(),
     43             ..State::default()
     44         })))
     45     }
     46 
     47     pub(crate) fn register(
     48         &self,
     49         targets: impl Iterator<Item = String>,
     50         budget: Arc<FetchBudget>,
     51     ) -> Option<FetchRegistration> {
     52         let mut state = self.0.lock().ok()?;
     53         if state.active.len() == MAX_ACTIVE_FETCHES {
     54             return None;
     55         }
     56         let targets: BTreeSet<_> = targets.collect();
     57         if targets
     58             .iter()
     59             .any(|target| !state.relays.contains_key(target))
     60         {
     61             return None;
     62         }
     63         let id = state.sequence.checked_add(1)?;
     64         state.sequence = id;
     65         state.active.insert(id, Registration { targets, budget });
     66         Some(FetchRegistration {
     67             registry: self.clone(),
     68             id,
     69             finished: false,
     70         })
     71     }
     72 
     73     pub(crate) fn connection(&self, relay: &str) -> Option<IngressConnection> {
     74         let mut state = self.0.lock().ok()?;
     75         let row = state.relays.get_mut(relay)?;
     76         row.generation = row.generation.checked_add(1)?;
     77         row.waker.wake();
     78         row.write_waker.wake();
     79         Some(IngressConnection {
     80             registry: self.clone(),
     81             relay: relay.to_owned(),
     82             generation: row.generation,
     83         })
     84     }
     85 
     86     fn admitted(&self, relay: &str, generation: u64, context: &Context<'_>, write: bool) -> bool {
     87         let Ok(state) = self.0.lock() else {
     88             return false;
     89         };
     90         let Some(row) = state.relays.get(relay) else {
     91             return false;
     92         };
     93         if generation <= row.invalidated || generation != row.generation {
     94             return false;
     95         }
     96         if write {
     97             row.write_waker.register(context.waker());
     98         } else {
     99             row.waker.register(context.waker());
    100         }
    101         true
    102     }
    103 
    104     fn charge(
    105         &self,
    106         relay: &str,
    107         generation: u64,
    108         bytes: usize,
    109         frames: usize,
    110         data: usize,
    111     ) -> bool {
    112         let Ok(mut state) = self.0.lock() else {
    113             return false;
    114         };
    115         let Some(row) = state.relays.get(relay) else {
    116             return false;
    117         };
    118         if generation <= row.invalidated || generation != row.generation {
    119             return false;
    120         }
    121         let mut denied = BTreeSet::new();
    122         for entry in state.active.values() {
    123             if entry.targets.contains(relay) && !entry.budget.wire(bytes, frames, data) {
    124                 denied.extend(entry.targets.iter().cloned());
    125             }
    126         }
    127         for target in &denied {
    128             if let Some(row) = state.relays.get_mut(target) {
    129                 row.invalidated = row.generation;
    130                 row.waker.wake();
    131                 row.write_waker.wake();
    132             }
    133         }
    134         denied.is_empty()
    135     }
    136 
    137     fn remove(&self, id: u64, cancel: bool) -> bool {
    138         let Ok(mut state) = self.0.lock() else {
    139             return false;
    140         };
    141         let Some(entry) = state.active.remove(&id) else {
    142             return false;
    143         };
    144         // Charge and completion use this same mutex. No ingress can exhaust
    145         // this registration between the budget check and its removal.
    146         let complete = !cancel && !entry.budget.exhausted();
    147         if !complete {
    148             for target in entry.targets {
    149                 if let Some(row) = state.relays.get_mut(&target) {
    150                     row.invalidated = row.generation;
    151                     row.waker.wake();
    152                     row.write_waker.wake();
    153                 }
    154             }
    155         }
    156         complete
    157     }
    158 }
    159 
    160 pub(crate) struct FetchRegistration {
    161     registry: IngressRegistry,
    162     id: u64,
    163     finished: bool,
    164 }
    165 
    166 impl FetchRegistration {
    167     pub(crate) fn finish(mut self) -> bool {
    168         let complete = self.registry.remove(self.id, false);
    169         self.finished = true;
    170         complete
    171     }
    172 }
    173 
    174 impl Drop for FetchRegistration {
    175     fn drop(&mut self) {
    176         if !self.finished {
    177             self.registry.remove(self.id, true);
    178         }
    179     }
    180 }
    181 
    182 #[derive(Clone, Debug)]
    183 pub(crate) struct IngressConnection {
    184     registry: IngressRegistry,
    185     relay: String,
    186     generation: u64,
    187 }
    188 
    189 impl IngressConnection {
    190     pub(crate) fn admitted(&self, context: &Context<'_>) -> bool {
    191         self.registry
    192             .admitted(&self.relay, self.generation, context, false)
    193     }
    194 
    195     pub(crate) fn admitted_write(&self, context: &Context<'_>) -> bool {
    196         self.registry
    197             .admitted(&self.relay, self.generation, context, true)
    198     }
    199 
    200     pub(crate) fn charge(&self, bytes: usize, frames: usize, data: usize) -> bool {
    201         self.registry
    202             .charge(&self.relay, self.generation, bytes, frames, data)
    203     }
    204 }
    205 
    206 #[cfg(test)]
    207 mod tests {
    208     use super::*;
    209     use crate::source::budget::{MAX_FETCH_BYTES, MAX_FETCH_EVENTS, MAX_FETCH_NOTIFICATIONS};
    210 
    211     fn registry() -> IngressRegistry {
    212         IngressRegistry::new(["one".to_owned(), "two".to_owned()].into_iter())
    213     }
    214 
    215     fn register(
    216         registry: &IngressRegistry,
    217         targets: &[&str],
    218         budget: &Arc<FetchBudget>,
    219     ) -> FetchRegistration {
    220         registry
    221             .register(
    222                 targets.iter().map(|target| (*target).to_owned()),
    223                 Arc::clone(budget),
    224             )
    225             .unwrap()
    226     }
    227 
    228     fn admitted(connection: &IngressConnection) -> bool {
    229         connection.admitted(&Context::from_waker(futures::task::noop_waker_ref()))
    230     }
    231 
    232     #[test]
    233     fn cancellation_closes_the_exact_generation_and_a_new_fetch_can_reconnect() {
    234         let registry = registry();
    235         let budget = Arc::new(FetchBudget::default());
    236         let registration = register(&registry, &["one"], &budget);
    237         let first = registry.connection("one").unwrap();
    238         assert!(admitted(&first));
    239         drop(registration);
    240         assert!(!admitted(&first));
    241         assert!(!first.charge(1, 1, 1));
    242         let fresh = register(&registry, &["one"], &Arc::new(FetchBudget::default()));
    243         let second = registry.connection("one").unwrap();
    244         assert!(admitted(&second));
    245         assert!(second.charge(1, 1, 1));
    246         assert!(!admitted(&first));
    247         fresh.finish();
    248         assert!(admitted(&second));
    249         assert!(second.charge(MAX_FETCH_BYTES + 1, 0, 0));
    250     }
    251 
    252     #[test]
    253     fn all_batches_and_reconnections_share_one_monotonic_budget() {
    254         let registry = registry();
    255         let budget = Arc::new(FetchBudget::default());
    256         let registration = register(&registry, &["one", "two"], &budget);
    257         let first = registry.connection("one").unwrap();
    258         assert!(first.charge(MAX_FETCH_BYTES / 2, 1, 1));
    259         let retried = registry.connection("one").unwrap();
    260         assert!(!admitted(&first));
    261         assert!(!first.charge(1, 0, 0));
    262         assert!(retried.charge(MAX_FETCH_BYTES / 2, 1, 1));
    263         let second = registry.connection("two").unwrap();
    264         assert!(!second.charge(1, 0, 0));
    265         assert!(budget.exhausted());
    266         assert!(!admitted(&retried));
    267         assert!(!admitted(&second));
    268         drop(registration);
    269         assert!(!admitted(&retried));
    270     }
    271 
    272     #[test]
    273     fn overlapping_fetches_each_charge_shared_traffic_and_unselected_relays_do_not() {
    274         let registry = registry();
    275         let first_budget = Arc::new(FetchBudget::default());
    276         let second_budget = Arc::new(FetchBudget::default());
    277         let first = register(&registry, &["one"], &first_budget);
    278         let second = register(&registry, &["one"], &second_budget);
    279         let other = registry.connection("two").unwrap();
    280         assert!(other.charge(
    281             MAX_FETCH_BYTES + 1,
    282             MAX_FETCH_NOTIFICATIONS + 1,
    283             MAX_FETCH_EVENTS + 1
    284         ));
    285         let connection = registry.connection("one").unwrap();
    286         assert!(connection.charge(MAX_FETCH_BYTES, 0, 0));
    287         first.finish();
    288         assert!(!connection.charge(1, 0, 0));
    289         assert!(!first_budget.exhausted());
    290         assert!(second_budget.exhausted());
    291         assert!(admitted(&other));
    292         drop(second);
    293     }
    294 
    295     #[test]
    296     fn registration_capacity_unknown_targets_and_counter_overflow_fail_closed() {
    297         let registry = registry();
    298         let budget = Arc::new(FetchBudget::default());
    299         assert!(
    300             registry
    301                 .register(["unknown".to_owned()].into_iter(), Arc::clone(&budget))
    302                 .is_none()
    303         );
    304         assert!(registry.connection("unknown").is_none());
    305         let mut registrations = (0..MAX_ACTIVE_FETCHES)
    306             .map(|_| register(&registry, &["one"], &budget))
    307             .collect::<Vec<_>>();
    308         assert!(
    309             registry
    310                 .register(["one".to_owned()].into_iter(), Arc::clone(&budget))
    311                 .is_none()
    312         );
    313         registrations.pop().unwrap().finish();
    314         register(&registry, &["one"], &budget).finish();
    315         for registration in registrations {
    316             registration.finish();
    317         }
    318         registry.0.lock().unwrap().sequence = u64::MAX;
    319         assert!(
    320             registry
    321                 .register(["one".to_owned()].into_iter(), budget)
    322                 .is_none()
    323         );
    324         registry
    325             .0
    326             .lock()
    327             .unwrap()
    328             .relays
    329             .get_mut("one")
    330             .unwrap()
    331             .generation = u64::MAX;
    332         assert!(registry.connection("one").is_none());
    333     }
    334 
    335     #[test]
    336     fn cancel_notifies_the_pending_reader() {
    337         struct Wake(std::sync::atomic::AtomicBool);
    338         impl futures::task::ArcWake for Wake {
    339             fn wake_by_ref(arc_self: &Arc<Self>) {
    340                 arc_self.0.store(true, std::sync::atomic::Ordering::SeqCst);
    341             }
    342         }
    343         let registry = registry();
    344         let registration = register(&registry, &["one"], &Arc::new(FetchBudget::default()));
    345         let connection = registry.connection("one").unwrap();
    346         let wake = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false)));
    347         let waker = futures::task::waker(Arc::clone(&wake));
    348         assert!(connection.admitted(&Context::from_waker(&waker)));
    349         drop(registration);
    350         assert!(wake.0.load(std::sync::atomic::Ordering::SeqCst));
    351         assert!(!admitted(&connection));
    352     }
    353 
    354     #[test]
    355     fn a_stale_generation_cannot_steal_live_reader_or_writer_wakeups() {
    356         struct Wake(std::sync::atomic::AtomicBool);
    357         impl futures::task::ArcWake for Wake {
    358             fn wake_by_ref(arc_self: &Arc<Self>) {
    359                 arc_self.0.store(true, std::sync::atomic::Ordering::SeqCst);
    360             }
    361         }
    362         let registry = registry();
    363         let registration = register(&registry, &["one"], &Arc::new(FetchBudget::default()));
    364         let stale = registry.connection("one").unwrap();
    365         let current = registry.connection("one").unwrap();
    366         let read = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false)));
    367         let write = Arc::new(Wake(std::sync::atomic::AtomicBool::new(false)));
    368         let read_waker = futures::task::waker(Arc::clone(&read));
    369         let write_waker = futures::task::waker(Arc::clone(&write));
    370         assert!(current.admitted(&Context::from_waker(&read_waker)));
    371         assert!(current.admitted_write(&Context::from_waker(&write_waker)));
    372         assert!(!stale.admitted(&Context::from_waker(futures::task::noop_waker_ref())));
    373         assert!(!stale.admitted_write(&Context::from_waker(futures::task::noop_waker_ref())));
    374         drop(registration);
    375         assert!(read.0.load(std::sync::atomic::Ordering::SeqCst));
    376         assert!(write.0.load(std::sync::atomic::Ordering::SeqCst));
    377     }
    378 
    379     #[test]
    380     fn finalization_and_ingress_have_one_ordered_completion_boundary() {
    381         for ingress_first in [false, true] {
    382             let registry = registry();
    383             let budget = Arc::new(FetchBudget::default());
    384             let registration = register(&registry, &["one"], &budget);
    385             let connection = registry.connection("one").unwrap();
    386             assert!(connection.charge(MAX_FETCH_BYTES, 0, 0));
    387             if ingress_first {
    388                 assert!(!connection.charge(1, 0, 0));
    389                 assert!(!registration.finish());
    390                 assert!(budget.exhausted());
    391                 assert!(!admitted(&connection));
    392             } else {
    393                 assert!(registration.finish());
    394                 assert!(connection.charge(1, 0, 0));
    395                 assert!(!budget.exhausted());
    396                 assert!(admitted(&connection));
    397             }
    398             assert!(registry.0.lock().unwrap().active.is_empty());
    399         }
    400     }
    401 
    402     #[test]
    403     fn racing_completion_and_last_byte_cannot_both_claim_admission() {
    404         for _ in 0..32 {
    405             let registry = registry();
    406             let budget = Arc::new(FetchBudget::default());
    407             let registration = register(&registry, &["one"], &budget);
    408             let connection = registry.connection("one").unwrap();
    409             assert!(connection.charge(MAX_FETCH_BYTES, 0, 0));
    410             let barrier = std::sync::Barrier::new(2);
    411             let (complete, admitted_after_limit) = std::thread::scope(|scope| {
    412                 let finish = scope.spawn(|| {
    413                     barrier.wait();
    414                     registration.finish()
    415                 });
    416                 barrier.wait();
    417                 let admitted_after_limit = connection.charge(1, 0, 0);
    418                 (finish.join().unwrap(), admitted_after_limit)
    419             });
    420             assert_eq!(complete, admitted_after_limit);
    421             assert_eq!(budget.exhausted(), !complete);
    422         }
    423     }
    424 
    425     #[test]
    426     fn missing_or_poisoned_registration_cannot_finalize_successfully() {
    427         let registry = registry();
    428         let budget = Arc::new(FetchBudget::default());
    429         let missing = register(&registry, &["one"], &budget);
    430         registry.0.lock().unwrap().active.remove(&missing.id);
    431         assert!(!missing.finish());
    432         let poisoned = register(&registry, &["one"], &budget);
    433         let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
    434             let _guard = registry.0.lock().unwrap();
    435             panic!("poison the registration lock");
    436         }));
    437         assert!(!poisoned.finish());
    438     }
    439 }