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(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["one"], &first_budget); 278 let second = register(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["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(®istry, &["one"], &budget); 430 registry.0.lock().unwrap().active.remove(&missing.id); 431 assert!(!missing.finish()); 432 let poisoned = register(®istry, &["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 }