diff --git a/src/quota_state.rs b/src/quota_state.rs index 1516b20..3d0dad6 100644 --- a/src/quota_state.rs +++ b/src/quota_state.rs @@ -104,14 +104,13 @@ impl QuotaStateOwner { used_bytes: 0, last_reset_epoch_secs, }; + let reset_target = self.store.current_or_legacy_handle(user); let state = self.state_for_users(configured_users, Some((user, prospective.clone()))); let path = self.path.clone(); - let store = Arc::clone(&self.store); - let user = user.to_string(); let task = tokio::task::spawn_blocking(move || { let _guard = guard; write_state_file_blocking(&path, &state)?; - Ok(store.reset(&user, last_reset_epoch_secs)) + Ok(reset_target.reset(last_reset_epoch_secs)) }); wait_for_blocking_io(task).await } diff --git a/src/stats/quota_store.rs b/src/stats/quota_store.rs index bf1d069..73446e3 100644 --- a/src/stats/quota_store.rs +++ b/src/stats/quota_store.rs @@ -213,12 +213,7 @@ impl QuotaStore { } pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot { - let state = self.current_or_legacy_handle(user); - state.counters.replace(0, now_epoch_secs); - UserQuotaSnapshot { - used_bytes: 0, - last_reset_epoch_secs: now_epoch_secs, - } + self.current_or_legacy_handle(user).reset(now_epoch_secs) } pub(crate) fn remove(&self, user: &str) { @@ -359,6 +354,15 @@ impl UserQuotaHandle { ) -> Result { self.counters.try_reserve(bytes, limit) } + + /// Resets only the quota incarnation captured by this handle. + pub(crate) fn reset(&self, now_epoch_secs: u64) -> UserQuotaSnapshot { + self.counters.replace(0, now_epoch_secs); + UserQuotaSnapshot { + used_bytes: 0, + last_reset_epoch_secs: now_epoch_secs, + } + } } impl QuotaReservation { @@ -497,6 +501,22 @@ mod tests { assert_eq!(current.used(), 40); } + #[test] + fn captured_reset_handle_cannot_reset_a_new_incarnation() { + let store = QuotaStore::default(); + store.activate_fresh("alice", 1); + let reset_target = store.handle_exact("alice", 1).unwrap(); + reset_target.charge(40); + store.advance_preserving_usage("alice", 2); + let current = store.handle_exact("alice", 2).unwrap(); + current.charge(20); + + reset_target.reset(7); + + assert_eq!(reset_target.used(), 0); + assert_eq!(current.used(), 60); + } + #[test] fn stale_retirement_cannot_remove_newer_quota_owner() { let store = QuotaStore::default(); diff --git a/src/transport/middle_proxy/health/family.rs b/src/transport/middle_proxy/health/family.rs index 0bb5d72..a23a02d 100644 --- a/src/transport/middle_proxy/health/family.rs +++ b/src/transport/middle_proxy/health/family.rs @@ -26,6 +26,7 @@ pub(super) async fn check_family( let mut dc_endpoints = HashMap::>::new(); let endpoint_snapshot = pool.endpoint_snapshot.load(); + let endpoint_revision = endpoint_snapshot.revision; let map_guard = match family { IpFamily::V4 => &endpoint_snapshot.map_v4, IpFamily::V6 => &endpoint_snapshot.map_v6, @@ -253,6 +254,7 @@ pub(super) async fn check_family( dc, family, generation: pool.current_generation(), + endpoint_revision, contour: WriterContour::Active, }) .await diff --git a/src/transport/middle_proxy/pool.rs b/src/transport/middle_proxy/pool.rs index d34fda7..6d58df2 100644 --- a/src/transport/middle_proxy/pool.rs +++ b/src/transport/middle_proxy/pool.rs @@ -37,6 +37,8 @@ pub(super) struct RefillTargetKey { pub family: IpFamily, /// Generation that retains publication authority. pub generation: u64, + /// Endpoint snapshot revision targeted by this refill producer. + pub endpoint_revision: u64, /// Lifecycle contour that the replacement must preserve. pub contour: WriterContour, } @@ -310,6 +312,7 @@ pub(super) struct ReinitStatusSnapshot { pub(super) pending_hardswap_generation: u64, pub(super) pending_hardswap_started_at_epoch_secs: u64, pub(super) pending_hardswap_map_hash: u64, + pub(super) pending_hardswap_endpoint_revision: u64, pub(super) inflight: usize, } diff --git a/src/transport/middle_proxy/pool/construction.rs b/src/transport/middle_proxy/pool/construction.rs index e66ddc1..68eaeec 100644 --- a/src/transport/middle_proxy/pool/construction.rs +++ b/src/transport/middle_proxy/pool/construction.rs @@ -143,6 +143,7 @@ impl MePool { pending_hardswap_generation: 0, pending_hardswap_started_at_epoch_secs: 0, pending_hardswap_map_hash: 0, + pending_hardswap_endpoint_revision: 0, inflight: 0, }; stats.set_me_writer_byte_budget_limit_bytes(me_writer_byte_budget_bytes); diff --git a/src/transport/middle_proxy/pool/routing.rs b/src/transport/middle_proxy/pool/routing.rs index c183404..623cd54 100644 --- a/src/transport/middle_proxy/pool/routing.rs +++ b/src/transport/middle_proxy/pool/routing.rs @@ -1,5 +1,37 @@ use super::*; +impl EndpointSnapshot { + /// Returns authoritative configured endpoints for one exact DC and address family. + pub(in crate::transport::middle_proxy) fn endpoints_for_dc_family( + &self, + dc: i32, + family: IpFamily, + ) -> &[(IpAddr, u16)] { + match family { + IpFamily::V4 => self.map_v4.get(&dc), + IpFamily::V6 => self.map_v6.get(&dc), + } + .map(Vec::as_slice) + .unwrap_or_default() + } + + /// Checks control-plane membership without applying data-plane family preference. + pub(in crate::transport::middle_proxy) fn contains_dc_endpoint( + &self, + dc: i32, + endpoint: SocketAddr, + ) -> bool { + let family = if endpoint.is_ipv4() { + IpFamily::V4 + } else { + IpFamily::V6 + }; + self.endpoints_for_dc_family(dc, family) + .iter() + .any(|(ip, port)| *ip == endpoint.ip() && *port == endpoint.port()) + } +} + impl MePool { pub(in crate::transport::middle_proxy) fn single_endpoint_outage_mode_enabled(&self) -> bool { self.single_endpoint_runtime diff --git a/src/transport/middle_proxy/pool/writer_admission.rs b/src/transport/middle_proxy/pool/writer_admission.rs index b11673c..3c20008 100644 --- a/src/transport/middle_proxy/pool/writer_admission.rs +++ b/src/transport/middle_proxy/pool/writer_admission.rs @@ -65,7 +65,7 @@ impl MePool { pub(in crate::transport::middle_proxy) async fn active_coverage_required_total(&self) -> usize { let now_epoch_secs = Self::now_epoch_secs(); let mut required_total = 0usize; - let endpoint_snapshot = self.endpoint_snapshot.load(); + let endpoint_snapshot = self.endpoint_snapshot.load_full(); if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { for addrs in endpoint_snapshot.map_v4.values() { @@ -100,11 +100,21 @@ impl MePool { contour: WriterContour, intent: WriterOpenIntent, writer_dc: i32, - family: IpFamily, + target_addr: SocketAddr, ) -> bool { if intent == WriterOpenIntent::Replacement { return true; } + let family = if target_addr.is_ipv4() { + IpFamily::V4 + } else { + IpFamily::V6 + }; + let endpoint_snapshot = self.endpoint_snapshot.load_full(); + let endpoints = endpoint_snapshot.endpoints_for_dc_family(writer_dc, family); + if !endpoint_snapshot.contains_dc_endpoint(writer_dc, target_addr) { + return false; + } let (active_writers, warm_writers, _) = self.non_draining_writer_counts_by_contour().await; let live = match contour { WriterContour::Active => active_writers, @@ -123,13 +133,7 @@ impl MePool { return false; } - let endpoint_snapshot = self.endpoint_snapshot.load(); - let endpoint_count = match family { - IpFamily::V4 => endpoint_snapshot.map_v4.get(&writer_dc), - IpFamily::V6 => endpoint_snapshot.map_v6.get(&writer_dc), - } - .map(Vec::len) - .unwrap_or(0); + let endpoint_count = endpoints.len(); if endpoint_count == 0 { return false; } @@ -150,6 +154,7 @@ impl MePool { && writer.generation == generation && WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) == contour && writer.addr.is_ipv4() == (family == IpFamily::V4) + && endpoint_snapshot.contains_dc_endpoint(writer_dc, writer.addr) }) .count() }; @@ -157,7 +162,7 @@ impl MePool { return true; } - live < self.active_coverage_required_total().await + false } /// Reserves bounded transient capacity for a writer open attempt. @@ -166,7 +171,7 @@ impl MePool { contour: WriterContour, intent: WriterOpenIntent, writer_dc: i32, - family: IpFamily, + target_addr: SocketAddr, ) -> Option> { let counter = match contour { WriterContour::Active => &self.writer_connect_active_reserved, @@ -221,7 +226,7 @@ impl MePool { loop { if !self - .can_open_writer_for_contour(contour, intent, writer_dc, family) + .can_open_writer_for_contour(contour, intent, writer_dc, target_addr) .await { return None; diff --git a/src/transport/middle_proxy/pool_refill.rs b/src/transport/middle_proxy/pool_refill.rs index 60ec7e6..f34cd44 100644 --- a/src/transport/middle_proxy/pool_refill.rs +++ b/src/transport/middle_proxy/pool_refill.rs @@ -26,6 +26,13 @@ enum RefillOutcome { Obsolete, } +fn refill_open_intent(contour: WriterContour) -> WriterOpenIntent { + match contour { + WriterContour::Active | WriterContour::Warm => WriterOpenIntent::Coverage, + WriterContour::Draining => WriterOpenIntent::Normal, + } +} + struct RefillRunGuard { pool: Arc, key: RefillTargetKey, @@ -293,18 +300,15 @@ impl MePool { && target.generation == status.pending_hardswap_generation, WriterContour::Draining => false, }; + let pending_revision_matches = target.contour != WriterContour::Warm + || target.endpoint_revision == status.pending_hardswap_endpoint_revision; + let endpoint_snapshot = self.endpoint_snapshot.load(); role_is_authoritative - && self - .endpoint_snapshot - .load() - .preferred_endpoints_by_dc - .get(&target.dc) - .is_some_and(|endpoints| { - endpoints.iter().any(|endpoint| match target.family { - IpFamily::V4 => endpoint.is_ipv4(), - IpFamily::V6 => endpoint.is_ipv6(), - }) - }) + && pending_revision_matches + && target.endpoint_revision == endpoint_snapshot.revision + && !endpoint_snapshot + .endpoints_for_dc_family(target.dc, target.family) + .is_empty() } async fn refill_writer_after_loss( @@ -315,11 +319,7 @@ impl MePool { if !self.refill_target_is_authoritative(target) { return RefillOutcome::Obsolete; } - let open_intent = if target.contour == WriterContour::Active { - WriterOpenIntent::Coverage - } else { - WriterOpenIntent::Normal - }; + let open_intent = refill_open_intent(target.contour); let fast_retries = self.reconnect_runtime.me_reconnect_fast_retry_count.max(1); let mut total_attempts = 0u32; let same_endpoint_quarantined = self.is_endpoint_quarantined(addr).await; @@ -455,6 +455,7 @@ impl MePool { dc: role.dc, family: role.family, generation: role.generation, + endpoint_revision: self.endpoint_snapshot.load().revision, contour: role.contour, }; if !self.refill_target_is_authoritative(target) { @@ -517,3 +518,16 @@ impl MePool { }); } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn warm_refill_preserves_required_floor_coverage() { + assert_eq!( + refill_open_intent(WriterContour::Warm), + WriterOpenIntent::Coverage + ); + } +} diff --git a/src/transport/middle_proxy/pool_reinit.rs b/src/transport/middle_proxy/pool_reinit.rs index 177561c..c07b228 100644 --- a/src/transport/middle_proxy/pool_reinit.rs +++ b/src/transport/middle_proxy/pool_reinit.rs @@ -24,6 +24,8 @@ mod coordination; // Generation warmup and stale-writer reconciliation. mod reconcile; +#[cfg(test)] +mod dual_family_tests; #[cfg(test)] mod tests; const ME_HARDSWAP_PENDING_TTL_SECS: u64 = 1800; @@ -115,6 +117,8 @@ fn publish_reinit_state(reinit: &ReinitCore, state: &ReinitCoordinatorState) { pending_hardswap_started_at_epoch_secs: pending .map_or(0, |value| value.started_at_epoch_secs), pending_hardswap_map_hash: pending.map_or(0, |value| value.map_hash), + pending_hardswap_endpoint_revision: pending + .map_or(0, |value| value.endpoint_revision), inflight: state.attempts.len(), }; reinit diff --git a/src/transport/middle_proxy/pool_reinit/coordination.rs b/src/transport/middle_proxy/pool_reinit/coordination.rs index 0b270eb..1b857aa 100644 --- a/src/transport/middle_proxy/pool_reinit/coordination.rs +++ b/src/transport/middle_proxy/pool_reinit/coordination.rs @@ -388,7 +388,7 @@ impl MePool { } /// Projects desired per-DC endpoint sets from one immutable endpoint revision. - pub(super) fn desired_dc_endpoints_from_snapshot( + pub(in crate::transport::middle_proxy) fn desired_dc_endpoints_from_snapshot( &self, endpoint_snapshot: &EndpointSnapshot, ) -> HashMap> { @@ -424,7 +424,6 @@ impl MePool { let active_generation = state.active_generation; let pending_generation = state.pending.map(|pending| pending.generation); let endpoint_snapshot = self.endpoint_snapshot.load(); - let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let now_epoch_secs = Self::now_epoch_secs(); let mut changed = 0usize; @@ -440,9 +439,7 @@ impl MePool { }; let endpoint_is_current = self .family_enabled_for_drain_coverage(family, now_epoch_secs) - && preferred - .get(&writer.writer_dc) - .is_some_and(|endpoints| endpoints.contains(&writer.addr)); + && endpoint_snapshot.contains_dc_endpoint(writer.writer_dc, writer.addr); if contour == WriterContour::Warm && writer.generation == active_generation && endpoint_is_current diff --git a/src/transport/middle_proxy/pool_reinit/dual_family_tests.rs b/src/transport/middle_proxy/pool_reinit/dual_family_tests.rs new file mode 100644 index 0000000..1e4cd0a --- /dev/null +++ b/src/transport/middle_proxy/pool_reinit/dual_family_tests.rs @@ -0,0 +1,129 @@ +use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; +use std::sync::atomic::Ordering; + +use super::tests::{insert_writer, insert_writer_floor}; +use crate::network::probe::NetworkDecision; +use crate::transport::middle_proxy::pool::WriterContour; +use crate::transport::middle_proxy::pool_writer_security_tests::make_pool_with_decision; + +#[tokio::test] +async fn nonpreferred_family_warm_generation_retains_authority_and_commits() { + let pool = make_pool_with_decision(NetworkDecision { + ipv4_me: true, + ipv6_me: true, + effective_prefer: 4, + effective_multipath: false, + ..NetworkDecision::default() + }) + .await; + let v4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let v6 = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443); + pool.update_proxy_maps( + HashMap::from([(2, vec![(v4.ip(), v4.port())])]), + Some(HashMap::from([(2, vec![(v6.ip(), v6.port())])])), + ) + .await; + let desired = pool.desired_dc_endpoints().await; + let map_hash = super::MePool::desired_map_hash(&desired); + let endpoint_revision = pool.endpoint_snapshot.load().revision; + let reservation = pool + .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) + .expect("current endpoint revision must admit the hardswap"); + assert!(pool.hardswap_warmup_is_authoritative( + reservation.attempt.generation, + map_hash, + endpoint_revision, + )); + let generation = reservation.attempt.generation; + let fresh_v4 = insert_writer_floor(&pool, 10, 2, v4, generation, WriterContour::Warm).await; + let fresh_v6 = insert_writer_floor(&pool, 20, 2, v6, generation, WriterContour::Warm).await; + let fresh_v4_media = + insert_writer_floor(&pool, 30, -2, v4, generation, WriterContour::Warm).await; + let fresh_v6_media = + insert_writer_floor(&pool, 40, -2, v6, generation, WriterContour::Warm).await; + + assert_eq!(pool.reconcile_writer_generation_roles().await, 0); + for writer in fresh_v4 + .iter() + .chain(&fresh_v6) + .chain(&fresh_v4_media) + .chain(&fresh_v6_media) + { + assert!(!writer.draining.load(Ordering::Acquire)); + } + + let outcome = pool + .commit_reinit_attempt(&reservation.attempt, &desired, 1.0) + .await + .expect("full dual-family floor must commit without multipath selection"); + + assert!(outcome.missing_groups.is_empty()); + assert_eq!(pool.current_generation(), generation); +} + +#[tokio::test] +async fn endpoint_revision_fences_pending_hardswap_after_map_aba() { + let pool = make_pool_with_decision(NetworkDecision { + ipv4_me: true, + effective_prefer: 4, + ..NetworkDecision::default() + }) + .await; + let endpoint_a = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 10)), 443); + let endpoint_b = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 11)), 443); + pool.update_proxy_maps( + HashMap::from([(2, vec![(endpoint_a.ip(), endpoint_a.port())])]), + None, + ) + .await; + let desired_a = pool.desired_dc_endpoints().await; + let map_hash = super::MePool::desired_map_hash(&desired_a); + let endpoint_revision = pool.endpoint_snapshot.load().revision; + let reservation = pool + .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) + .expect("current endpoint revision must admit the hardswap"); + assert!(pool.hardswap_warmup_is_authoritative( + reservation.attempt.generation, + map_hash, + endpoint_revision, + )); + let warm = insert_writer( + &pool, + 100, + 2, + endpoint_a, + reservation.attempt.generation, + WriterContour::Warm, + ) + .await; + + pool.update_proxy_maps( + HashMap::from([(2, vec![(endpoint_b.ip(), endpoint_b.port())])]), + None, + ) + .await; + assert!(!pool.hardswap_warmup_is_authoritative( + reservation.attempt.generation, + map_hash, + endpoint_revision, + )); + pool.update_proxy_maps( + HashMap::from([(2, vec![(endpoint_a.ip(), endpoint_a.port())])]), + None, + ) + .await; + assert!(!pool.hardswap_warmup_is_authoritative( + reservation.attempt.generation, + map_hash, + endpoint_revision, + )); + + let snapshot = pool.api_hardswap_snapshot().await; + + assert!(snapshot.pending); + assert_eq!(snapshot.pending_map_current, Some(false)); + assert_eq!(snapshot.pending_writers_current, 0); + assert_eq!(snapshot.orphan_warm_writers_current, 1); + assert!(!warm.draining.load(Ordering::Acquire)); +} diff --git a/src/transport/middle_proxy/pool_reinit/reconcile.rs b/src/transport/middle_proxy/pool_reinit/reconcile.rs index 6ac3d1a..1efdbc9 100644 --- a/src/transport/middle_proxy/pool_reinit/reconcile.rs +++ b/src/transport/middle_proxy/pool_reinit/reconcile.rs @@ -1,10 +1,29 @@ use super::*; impl MePool { + /// Checks the exact pending-generation tuple before starting more warmup work. + pub(super) fn hardswap_warmup_is_authoritative( + &self, + generation: u64, + map_hash: u64, + endpoint_revision: u64, + ) -> bool { + let state = self.reinit.coordinator.lock(); + state.desired_map_hash == map_hash + && state.endpoint_revision == endpoint_revision + && state.pending.is_some_and(|pending| { + pending.generation == generation + && pending.map_hash == map_hash + && pending.endpoint_revision == endpoint_revision + }) + } + async fn warmup_generation_for_all_dcs( self: &Arc, rng: &SecureRandom, generation: u64, + map_hash: u64, + endpoint_revision: u64, desired_by_dc: &HashMap>, ) { let extra_passes = self @@ -15,6 +34,13 @@ impl MePool { let total_passes = 1 + extra_passes; for (dc, endpoints) in desired_by_dc { + if !self.hardswap_warmup_is_authoritative( + generation, + map_hash, + endpoint_revision, + ) { + return; + } for family in [IpFamily::V4, IpFamily::V6] { let family_endpoints = endpoints .iter() @@ -53,8 +79,22 @@ impl MePool { ); for attempt_idx in 0..missing { + if !self.hardswap_warmup_is_authoritative( + generation, + map_hash, + endpoint_revision, + ) { + return; + } let delay_ms = self.hardswap_warmup_connect_delay_ms(); tokio::time::sleep(Duration::from_millis(delay_ms)).await; + if !self.hardswap_warmup_is_authoritative( + generation, + map_hash, + endpoint_revision, + ) { + return; + } let connected = self .connect_endpoints_round_robin_with_generation_contour( @@ -100,6 +140,13 @@ impl MePool { } if pass_idx + 1 < total_passes { + if !self.hardswap_warmup_is_authoritative( + generation, + map_hash, + endpoint_revision, + ) { + return; + } let backoff_ms = self.hardswap_warmup_backoff_ms(pass_idx); debug!( dc = *dc, @@ -198,8 +245,14 @@ impl MePool { } if hardswap { - self.warmup_generation_for_all_dcs(rng, generation, &desired_by_dc) - .await; + self.warmup_generation_for_all_dcs( + rng, + generation, + desired_map_hash, + endpoint_snapshot.revision, + &desired_by_dc, + ) + .await; } else { self.reconcile_connections(rng).await; } diff --git a/src/transport/middle_proxy/pool_reinit/tests.rs b/src/transport/middle_proxy/pool_reinit/tests.rs index a0abe67..4abc8c0 100644 --- a/src/transport/middle_proxy/pool_reinit/tests.rs +++ b/src/transport/middle_proxy/pool_reinit/tests.rs @@ -26,7 +26,7 @@ fn addr_v6(segment: u16, port: u16) -> SocketAddr { ) } -async fn insert_writer( +pub(super) async fn insert_writer( pool: &Arc, writer_id: u64, writer_dc: i32, @@ -63,7 +63,7 @@ async fn insert_writer( writer } -async fn insert_writer_floor( +pub(super) async fn insert_writer_floor( pool: &Arc, first_writer_id: u64, writer_dc: i32, @@ -489,7 +489,7 @@ async fn generation_role_reconciliation_promotes_active_warm_and_drains_orphans( None, ) .await; - let desired_by_dc = HashMap::from([(1, HashSet::from([endpoint]))]); + let desired_by_dc = pool.desired_dc_endpoints().await; let map_hash = MePool::desired_map_hash(&desired_by_dc); let endpoint_revision = pool.endpoint_snapshot.load().revision; let reservation = pool diff --git a/src/transport/middle_proxy/pool_status.rs b/src/transport/middle_proxy/pool_status.rs index 32fd56f..42bf2e3 100644 --- a/src/transport/middle_proxy/pool_status.rs +++ b/src/transport/middle_proxy/pool_status.rs @@ -15,6 +15,9 @@ mod runtime_snapshot; // Hardswap ownership, coverage, and writer-replacement lifecycle state. mod hardswap_snapshot; pub(crate) use hardswap_snapshot::MeApiHardswapSnapshot; +#[cfg(test)] +#[path = "pool_status/tests.rs"] +mod status_invariant_tests; #[derive(Clone, Debug)] pub(crate) struct MeApiWriterStatusSnapshot { pub writer_id: u64, diff --git a/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs b/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs index 1185c12..1be4f2f 100644 --- a/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs +++ b/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs @@ -35,11 +35,15 @@ impl MePool { &self, reinit: &ReinitStatusSnapshot, ) -> MeApiHardswapSnapshot { - let desired_by_dc = self.desired_dc_endpoints().await; + let endpoint_snapshot = self.endpoint_snapshot.load_full(); + let desired_by_dc = self.desired_dc_endpoints_from_snapshot(&endpoint_snapshot); let desired_hash = Self::desired_map_hash(&desired_by_dc); let writers = self.writers.read().await; let pending_generation = reinit.pending_hardswap_generation; let pending = pending_generation != 0; + let pending_map_current = pending + && reinit.pending_hardswap_map_hash == desired_hash + && reinit.pending_hardswap_endpoint_revision == endpoint_snapshot.revision; let mut pending_writers_current = 0usize; let mut pending_writer_addrs = Vec::<(i32, SocketAddr)>::new(); let mut orphan_warm_writers_current = 0usize; @@ -49,10 +53,12 @@ impl MePool { continue; } let contour = WriterContour::from_u8(writer.contour.load(Ordering::Acquire)); - if contour == WriterContour::Warm && writer.generation != pending_generation { + if contour == WriterContour::Warm + && (!pending_map_current || writer.generation != pending_generation) + { orphan_warm_writers_current = orphan_warm_writers_current.saturating_add(1); } - if pending + if pending_map_current && writer.generation == pending_generation && contour == WriterContour::Warm && desired_by_dc @@ -64,7 +70,7 @@ impl MePool { } } - let pending_coverage = pending + let pending_coverage = pending_map_current .then(|| self.hardswap_coverage(&desired_by_dc, &pending_writer_addrs)); let pending_writer_deficit = pending_coverage .as_ref() @@ -85,8 +91,7 @@ impl MePool { pending_writers_current, pending_writer_deficit, pending_missing_dc_groups, - pending_map_current: pending - .then_some(reinit.pending_hardswap_map_hash == desired_hash), + pending_map_current: pending.then_some(pending_map_current), orphan_warm_writers_current, replacement_preparing_current, replacement_retiring_current, diff --git a/src/transport/middle_proxy/pool_status/status_snapshot.rs b/src/transport/middle_proxy/pool_status/status_snapshot.rs index a246b44..9379a63 100644 --- a/src/transport/middle_proxy/pool_status/status_snapshot.rs +++ b/src/transport/middle_proxy/pool_status/status_snapshot.rs @@ -62,7 +62,7 @@ impl MePool { let active_generation = self.reinit.status.load().active_generation; let writers = self.writers.read().await.clone(); - let mut live_writers_by_dc = HashMap::::new(); + let mut live_writers_by_group = HashMap::<(i16, bool), usize>::new(); for writer in writers.iter() { if writer.draining.load(Ordering::Relaxed) || writer.generation != active_generation @@ -72,19 +72,26 @@ impl MePool { continue; } if let Ok(dc) = i16::try_from(writer.writer_dc) { - *live_writers_by_dc.entry(dc).or_insert(0) += 1; + *live_writers_by_group + .entry((dc, writer.addr.is_ipv4())) + .or_insert(0) += 1; } } for (dc, endpoints) in endpoints_by_dc { - let endpoint_count = endpoints.len(); - if endpoint_count == 0 { - return false; - } - let required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false); - let alive = live_writers_by_dc.get(&dc).copied().unwrap_or(0); - if alive < required { - return false; + for (ipv4, endpoint_count) in endpoint_family_counts(&endpoints) { + if endpoint_count == 0 { + continue; + } + let required = + self.required_writers_for_dc_with_floor_mode(endpoint_count, false); + let alive = live_writers_by_group + .get(&(dc, ipv4)) + .copied() + .unwrap_or(0); + if alive < required { + return false; + } } } @@ -123,7 +130,14 @@ impl MePool { let required_writers = endpoints_by_dc .values() - .map(|endpoints| self.required_writers_for_dc_with_floor_mode(endpoints.len(), false)) + .map(|endpoints| { + endpoint_family_counts(endpoints) + .into_iter() + .map(|(_, count)| { + self.required_writers_for_dc_with_floor_mode(count, false) + }) + .sum::() + }) .sum(); let idle_since = self.registry.writer_idle_since_snapshot().await; @@ -227,39 +241,58 @@ impl MePool { .max(1); for (dc, endpoints) in endpoints_by_dc { let endpoint_count = endpoints.len(); + let family_counts = endpoint_family_counts(&endpoints); let dc_available_endpoints = endpoints .iter() .filter(|endpoint| live_writers_by_dc_endpoint.contains_key(&(dc, **endpoint))) .count(); - let base_required = self.required_writers_for_dc(endpoint_count); - let dc_required_writers = - self.required_writers_for_dc_with_floor_mode(endpoint_count, false); - let floor_min = if endpoint_count <= 1 { - (self - .floor_runtime - .me_adaptive_floor_min_writers_single_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - .min(base_required.max(1)) - } else { - (self - .floor_runtime - .me_adaptive_floor_min_writers_multi_endpoint - .load(Ordering::Relaxed) as usize) - .max(1) - .min(base_required.max(1)) - }; - let extra_per_core = if endpoint_count <= 1 { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_single_per_core - .load(Ordering::Relaxed) as usize - } else { - self.floor_runtime - .me_adaptive_floor_max_extra_writers_multi_per_core - .load(Ordering::Relaxed) as usize - }; - let floor_max = - base_required.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core)); + let base_required = family_counts + .iter() + .filter(|(_, count)| *count > 0) + .map(|(_, count)| self.required_writers_for_dc(*count)) + .sum::(); + let dc_required_writers = family_counts + .iter() + .filter(|(_, count)| *count > 0) + .map(|(_, count)| { + self.required_writers_for_dc_with_floor_mode(*count, false) + }) + .sum::(); + let floor_min = family_counts + .iter() + .filter(|(_, count)| *count > 0) + .map(|(_, count)| { + let family_base = self.required_writers_for_dc(*count); + let configured = if *count <= 1 { + self.floor_runtime + .me_adaptive_floor_min_writers_single_endpoint + .load(Ordering::Relaxed) as usize + } else { + self.floor_runtime + .me_adaptive_floor_min_writers_multi_endpoint + .load(Ordering::Relaxed) as usize + }; + configured.max(1).min(family_base.max(1)) + }) + .sum::(); + let floor_max = family_counts + .iter() + .filter(|(_, count)| *count > 0) + .map(|(_, count)| { + let family_base = self.required_writers_for_dc(*count); + let extra_per_core = if *count <= 1 { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_single_per_core + .load(Ordering::Relaxed) as usize + } else { + self.floor_runtime + .me_adaptive_floor_max_extra_writers_multi_per_core + .load(Ordering::Relaxed) as usize + }; + family_base + .saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core)) + }) + .sum::(); let floor_capped = matches!(floor_mode, MeFloorMode::Adaptive) && dc_required_writers < base_required; let dc_alive_writers = live_writers_by_dc.get(&dc).copied().unwrap_or(0); @@ -322,3 +355,8 @@ impl MePool { } } } + +fn endpoint_family_counts(endpoints: &BTreeSet) -> [(bool, usize); 2] { + let ipv4 = endpoints.iter().filter(|endpoint| endpoint.is_ipv4()).count(); + [(true, ipv4), (false, endpoints.len().saturating_sub(ipv4))] +} diff --git a/src/transport/middle_proxy/pool_status/tests.rs b/src/transport/middle_proxy/pool_status/tests.rs new file mode 100644 index 0000000..ead68a1 --- /dev/null +++ b/src/transport/middle_proxy/pool_status/tests.rs @@ -0,0 +1,84 @@ +use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64}; +use std::time::Instant; + +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +use crate::network::probe::NetworkDecision; +use crate::transport::middle_proxy::codec::WriterCommand; +use crate::transport::middle_proxy::pool::{MePool, MeWriter, WriterContour}; +use crate::transport::middle_proxy::pool_writer_security_tests::make_pool_with_decision; + +fn writer( + pool: &Arc, + id: u64, + dc: i32, + addr: SocketAddr, +) -> MeWriter { + let (tx, _rx) = mpsc::channel::(8); + MeWriter { + id, + addr, + source_ip: addr.ip(), + writer_dc: dc, + generation: pool.current_generation(), + contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())), + created_at: Instant::now(), + tx, + byte_budget: pool.new_writer_byte_budget(), + cancel: CancellationToken::new(), + degraded: Arc::new(AtomicBool::new(false)), + rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), + draining: Arc::new(AtomicBool::new(false)), + draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)), + drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)), + allow_drain_fallback: Arc::new(AtomicBool::new(false)), + } +} + +#[tokio::test] +async fn dual_family_status_reports_each_family_floor() { + let pool = make_pool_with_decision(NetworkDecision { + ipv4_me: true, + ipv6_me: true, + effective_prefer: 4, + effective_multipath: false, + ..NetworkDecision::default() + }) + .await; + let v4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let v6 = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443); + pool.update_proxy_maps( + HashMap::from([(2, vec![(v4.ip(), v4.port())])]), + Some(HashMap::from([(2, vec![(v6.ip(), v6.port())])])), + ) + .await; + let required_per_family = pool.required_writers_for_dc(1); + let mut writers = pool.writers.write().await; + for (group, dc) in [2, -2].into_iter().enumerate() { + for offset in 0..required_per_family { + writers.push(writer( + &pool, + (group as u64 * 100) + offset as u64, + dc, + v4, + )); + } + } + drop(writers); + + let snapshot = pool.api_status_snapshot().await; + + assert_eq!(snapshot.required_writers, required_per_family * 4); + assert_eq!(snapshot.alive_writers, required_per_family * 2); + assert_eq!(snapshot.coverage_pct, 50.0); + for dc in snapshot.dcs { + assert_eq!(dc.required_writers, required_per_family * 2); + assert_eq!(dc.alive_writers, required_per_family); + assert_eq!(dc.coverage_pct, 50.0); + } + assert!(!pool.admission_ready_full_floor().await); +} diff --git a/src/transport/middle_proxy/pool_writer/publication.rs b/src/transport/middle_proxy/pool_writer/publication.rs index b68b940..ef6fac2 100644 --- a/src/transport/middle_proxy/pool_writer/publication.rs +++ b/src/transport/middle_proxy/pool_writer/publication.rs @@ -43,9 +43,7 @@ impl MePool { let endpoint_is_current = self .endpoint_snapshot .load() - .preferred_endpoints_by_dc - .get(&writer.writer_dc) - .is_some_and(|endpoints| endpoints.contains(&writer.addr)); + .contains_dc_endpoint(writer.writer_dc, writer.addr); if !endpoint_is_current { return Err(ProxyError::Proxy( "ME writer target changed before publication".into(), @@ -95,21 +93,29 @@ impl MePool { return Ok(()); } let endpoint_snapshot = self.endpoint_snapshot.load(); - let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; - let Some(endpoints) = preferred.get(&writer.writer_dc) else { + if !endpoint_snapshot.contains_dc_endpoint(writer.writer_dc, writer.addr) { return Err(ProxyError::Proxy( "ME writer target changed before publication".into(), )); - }; - let required = match contour { - WriterContour::Active | WriterContour::Warm => self.required_writers_for_dc( - endpoints - .iter() - .filter(|endpoint| endpoint.is_ipv4() == writer.addr.is_ipv4()) - .count(), - ), - WriterContour::Draining => 0, - }; + } + if contour == WriterContour::Active && intent == WriterOpenIntent::Normal { + let current = writers + .iter() + .filter(|candidate| { + !candidate.draining.load(Ordering::Acquire) + && WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) + == WriterContour::Active + }) + .count(); + if current >= self.adaptive_floor_active_cap_configured_total() { + return Err(ProxyError::Proxy( + "ME active writer cap was reached before publication".into(), + )); + } + return Ok(()); + } + let endpoints = endpoint_snapshot.endpoints_for_dc_family(writer.writer_dc, family); + let required = self.required_writers_for_dc(endpoints.len()); let current = writers .iter() .filter(|candidate| { @@ -119,7 +125,8 @@ impl MePool { && WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) == contour && candidate.addr.is_ipv4() == writer.addr.is_ipv4() - && endpoints.contains(&candidate.addr) + && endpoint_snapshot + .contains_dc_endpoint(candidate.writer_dc, candidate.addr) }) .count(); if current >= required { diff --git a/src/transport/middle_proxy/pool_writer/replacement.rs b/src/transport/middle_proxy/pool_writer/replacement.rs index ff34bda..f1d05e9 100644 --- a/src/transport/middle_proxy/pool_writer/replacement.rs +++ b/src/transport/middle_proxy/pool_writer/replacement.rs @@ -86,7 +86,6 @@ impl MePool { )); } let endpoint_snapshot = self.endpoint_snapshot.load(); - let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let donor_count = writers .iter() .filter(|candidate| { @@ -95,9 +94,8 @@ impl MePool { && candidate.generation == expected_victim_role.generation && WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) == WriterContour::Active - && preferred - .get(&candidate.writer_dc) - .is_some_and(|endpoints| endpoints.contains(&candidate.addr)) + && endpoint_snapshot + .contains_dc_endpoint(candidate.writer_dc, candidate.addr) && (candidate.addr.is_ipv4() == matches!(expected_victim_role.family, crate::network::IpFamily::V4)) }) @@ -115,9 +113,8 @@ impl MePool { && candidate.generation == writer.generation && WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) == WriterContour::Active - && preferred - .get(&candidate.writer_dc) - .is_some_and(|endpoints| endpoints.contains(&candidate.addr)) + && endpoint_snapshot + .contains_dc_endpoint(candidate.writer_dc, candidate.addr) && (if candidate.addr.is_ipv4() { crate::network::IpFamily::V4 } else { @@ -223,11 +220,7 @@ mod tests { WriterContour::Active, WriterOpenIntent::Replacement, writer_dc, - if addr.is_ipv4() { - crate::network::IpFamily::V4 - } else { - crate::network::IpFamily::V6 - }, + addr, ) .await .expect("replacement open must be admitted"); diff --git a/src/transport/middle_proxy/pool_writer/runtime.rs b/src/transport/middle_proxy/pool_writer/runtime.rs index 412b2cd..4e5a5f2 100644 --- a/src/transport/middle_proxy/pool_writer/runtime.rs +++ b/src/transport/middle_proxy/pool_writer/runtime.rs @@ -96,11 +96,7 @@ impl MePool { contour, intent, writer_dc, - if addr.is_ipv4() { - crate::network::IpFamily::V4 - } else { - crate::network::IpFamily::V6 - }, + addr, ) .await else { diff --git a/src/transport/middle_proxy/send.rs b/src/transport/middle_proxy/send.rs index cec8795..3173534 100644 --- a/src/transport/middle_proxy/send.rs +++ b/src/transport/middle_proxy/send.rs @@ -1,26 +1,29 @@ #![allow(clippy::too_many_arguments)] -use std::cmp::Reverse; use std::net::SocketAddr; use std::sync::Arc; use std::sync::atomic::Ordering; use std::time::{Duration, Instant}; use tokio::sync::mpsc::error::TrySendError; -use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError, mpsc}; +use tokio::sync::{OwnedSemaphorePermit, TryAcquireError}; use tracing::{debug, warn}; use super::MePool; -use super::codec::{ProxyReqCommand, WriterBytePermit, WriterCommand}; +use super::codec::WriterCommand; use super::registry::{ConnMeta, WriterBindOutcome}; -use super::wire::{build_proxy_req_payload, proxy_req_payload_len}; -use crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES; -use crate::config::{MeRouteNoWriterMode, MeWriterPickMode}; +use super::wire::build_proxy_req_payload; +use crate::config::MeRouteNoWriterMode; use crate::error::{ProxyError, Result}; -use crate::stats::Stats; -use crate::stream::PooledBuffer; use rand::seq::SliceRandom; +use self::bound::BoundWriterSendOutcome; +use self::reservation::{ + LEGACY_PROXY_REQ_SOURCE_CAPACITY_OVERHEAD_BYTES, WriterByteReserveError, + WriterCommandReserveError, proxy_req_resident_permits, reserve_writer_bytes, + reserve_writer_command_slot, try_reserve_writer_bytes, writer_send_deadline, +}; + const IDLE_WRITER_PENALTY_MID_SECS: u64 = 45; const IDLE_WRITER_PENALTY_HIGH_SECS: u64 = 55; const HYBRID_GLOBAL_BURST_PERIOD_ROUNDS: u32 = 4; @@ -31,142 +34,15 @@ const PICK_PENALTY_WARM: u64 = 200; const PICK_PENALTY_DRAINING: u64 = 600; const PICK_PENALTY_STALE: u64 = 300; const PICK_PENALTY_DEGRADED: u64 = 250; -const RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES: usize = 27; -const LEGACY_PROXY_REQ_SOURCE_CAPACITY_OVERHEAD_BYTES: usize = 128; +// Send-path submodules isolate delivery, close handling, recovery, reservations, and selection. +mod bound; mod close; +mod pooled; mod recovery; +mod reservation; mod selection; -enum WriterCommandReserveError { - Closed, - TimedOut, -} - -enum WriterByteReserveError { - Closed, - TimedOut, -} - -fn proxy_tag_array(tag: Option<&[u8]>) -> Option<[u8; 16]> { - tag.and_then(|tag| <[u8; 16]>::try_from(tag).ok()) -} - -fn proxy_req_payload_from_command( - cmd: WriterCommand, -) -> Option<(PooledBuffer, OwnedSemaphorePermit)> { - match cmd { - WriterCommand::ProxyReq(command) => Some((command.payload, command._permit)), - _ => None, - } -} - -fn payload_permit_from_data_command(cmd: WriterCommand) -> Option { - match cmd { - WriterCommand::Data { _permit, .. } => _permit, - _ => None, - } -} - -async fn reserve_writer_command_slot( - tx: &mpsc::Sender, - deadline: Option, -) -> std::result::Result, WriterCommandReserveError> { - let reserve = tx.clone().reserve_owned(); - match deadline { - Some(deadline) => { - match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), reserve) - .await - { - Ok(Ok(permit)) => Ok(permit), - Ok(Err(_)) => Err(WriterCommandReserveError::Closed), - Err(_) => Err(WriterCommandReserveError::TimedOut), - } - } - None => reserve.await.map_err(|_| WriterCommandReserveError::Closed), - } -} - -fn writer_send_deadline(wait: Option) -> Option { - wait.map(|wait| Instant::now() + wait) -} - -fn writer_resident_permits( - source_capacity: usize, - encoded_payload_len: usize, -) -> Option<(u32, usize)> { - let resident_bytes = source_capacity - .checked_add(encoded_payload_len)? - .checked_add(RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES)?; - let permits = resident_bytes.div_ceil(ME_WRITER_BYTE_PERMIT_UNIT_BYTES); - let permits = u32::try_from(permits).ok()?; - let reserved_bytes = (permits as usize).checked_mul(ME_WRITER_BYTE_PERMIT_UNIT_BYTES)?; - Some(( - permits.max(1), - reserved_bytes.max(ME_WRITER_BYTE_PERMIT_UNIT_BYTES), - )) -} - -fn proxy_req_resident_permits( - source_capacity: usize, - data_len: usize, - proxy_tag: Option<&[u8]>, - proto_flags: u32, -) -> Option<(u32, usize)> { - writer_resident_permits( - source_capacity, - proxy_req_payload_len(data_len, proxy_tag, proto_flags), - ) -} - -fn try_reserve_writer_bytes( - byte_budget: &Arc, - permits: u32, - reserved_bytes: usize, - stats: &Arc, -) -> std::result::Result { - byte_budget - .clone() - .try_acquire_many_owned(permits) - .map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone())) -} - -async fn reserve_writer_bytes( - byte_budget: &Arc, - permits: u32, - reserved_bytes: usize, - deadline: Option, - stats: &Arc, -) -> std::result::Result { - match try_reserve_writer_bytes(byte_budget, permits, reserved_bytes, stats) { - Ok(permit) => return Ok(permit), - Err(TryAcquireError::Closed) => return Err(WriterByteReserveError::Closed), - Err(TryAcquireError::NoPermits) => { - stats.increment_me_writer_byte_budget_wait_total(); - } - } - - let acquire = byte_budget.clone().acquire_many_owned(permits); - match deadline { - Some(deadline) => { - match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), acquire) - .await - { - Ok(Ok(permit)) => Ok(WriterBytePermit::new(permit, reserved_bytes, stats.clone())), - Ok(Err(_)) => Err(WriterByteReserveError::Closed), - Err(_) => { - stats.increment_me_writer_byte_budget_timeout_total(); - Err(WriterByteReserveError::TimedOut) - } - } - } - None => acquire - .await - .map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone())) - .map_err(|_| WriterByteReserveError::Closed), - } -} - impl MePool { /// Send RPC_PROXY_REQ. `tag_override`: per-user ad_tag (from access.user_ad_tags); if None, uses pool default. /// `payload_permit` keeps optional client byte accounting alive until the writer consumes the command. @@ -243,78 +119,22 @@ impl MePool { let mut hybrid_wait_current = hybrid_wait_step; loop { - if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await - { - let deadline = - writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout); - let writer_permit = match reserve_writer_bytes( - ¤t.byte_budget, + match self + .try_send_bound_writer( + conn_id, + client_addr, + data, + proto_flags, + tag, writer_byte_permits, writer_reserved_bytes, - deadline, - &self.stats, + payload_permit, ) - .await - { - Ok(permit) => permit, - Err(WriterByteReserveError::TimedOut) => { - self.stats - .increment_me_writer_pick_full_total(self.writer_pick_mode()); - return Err(ProxyError::Proxy( - "ME writer byte budget full within blocking send timeout".into(), - )); - } - Err(WriterByteReserveError::Closed) => { - warn!( - writer_id = current.writer_id, - "ME writer byte budget closed" - ); - self.remove_writer_and_close_clients(current.writer_id) - .await; - continue; - } - }; - let (current_payload, _) = build_routed_payload(current_meta.our_addr); - let command = WriterCommand::Data { - payload: current_payload, - _permit: payload_permit.take(), - writer_permit, - }; - match current.tx.try_send(command) { - Ok(()) => { - self.note_hybrid_route_success(); - return Ok(()); - } - Err(TrySendError::Full(cmd)) => { - match reserve_writer_command_slot(¤t.tx, deadline).await { - Ok(permit) => { - permit.send(cmd); - self.note_hybrid_route_success(); - return Ok(()); - } - Err(WriterCommandReserveError::TimedOut) => { - self.stats - .increment_me_writer_pick_full_total(self.writer_pick_mode()); - return Err(ProxyError::Proxy( - "ME writer channel full within blocking send timeout".into(), - )); - } - Err(WriterCommandReserveError::Closed) => { - payload_permit = payload_permit_from_data_command(cmd); - } - } - warn!(writer_id = current.writer_id, "ME writer channel closed"); - self.remove_writer_and_close_clients(current.writer_id) - .await; - continue; - } - Err(TrySendError::Closed(cmd)) => { - payload_permit = payload_permit_from_data_command(cmd); - warn!(writer_id = current.writer_id, "ME writer channel closed"); - self.remove_writer_and_close_clients(current.writer_id) - .await; - continue; - } + .await? + { + BoundWriterSendOutcome::Sent => return Ok(()), + BoundWriterSendOutcome::Retry(retry_permit) => { + payload_permit = retry_permit; } } @@ -537,88 +357,9 @@ impl MePool { } hybrid_wait_current = hybrid_wait_step; let pick_mode = self.writer_pick_mode(); - let pick_sample_size = self.writer_pick_sample_size(); - let writer_ids: Vec = candidate_indices - .iter() - .map(|idx| writers_snapshot[*idx].id) - .collect(); - let writer_idle_since = self - .registry - .writer_idle_since_for_writer_ids(&writer_ids) + let ordered_candidate_indices = self + .ordered_candidate_indices(candidate_indices, &writers_snapshot, pick_mode) .await; - let now_epoch_secs = Self::now_epoch_secs(); - let start = self.rr.fetch_add(1, Ordering::Relaxed) as usize % candidate_indices.len(); - let ordered_candidate_indices = if pick_mode == MeWriterPickMode::P2c { - self.p2c_ordered_candidate_indices( - &candidate_indices, - &writers_snapshot, - &writer_idle_since, - now_epoch_secs, - start, - pick_sample_size, - ) - } else { - if self - .writer_selection_policy - .me_deterministic_writer_sort - .load(Ordering::Relaxed) - { - candidate_indices.sort_by(|lhs, rhs| { - let left = &writers_snapshot[*lhs]; - let right = &writers_snapshot[*rhs]; - let left_key = ( - self.writer_contour_rank_for_selection(left), - (left.generation < self.current_generation()) as usize, - left.degraded.load(Ordering::Relaxed) as usize, - self.writer_idle_rank_for_selection( - left, - &writer_idle_since, - now_epoch_secs, - ), - Reverse(left.tx.capacity()), - left.addr, - left.id, - ); - let right_key = ( - self.writer_contour_rank_for_selection(right), - (right.generation < self.current_generation()) as usize, - right.degraded.load(Ordering::Relaxed) as usize, - self.writer_idle_rank_for_selection( - right, - &writer_idle_since, - now_epoch_secs, - ), - Reverse(right.tx.capacity()), - right.addr, - right.id, - ); - left_key.cmp(&right_key) - }); - } else { - candidate_indices.sort_by_key(|idx| { - let w = &writers_snapshot[*idx]; - let degraded = w.degraded.load(Ordering::Relaxed); - let stale = (w.generation < self.current_generation()) as usize; - ( - self.writer_contour_rank_for_selection(w), - stale, - degraded as usize, - self.writer_idle_rank_for_selection( - w, - &writer_idle_since, - now_epoch_secs, - ), - Reverse(w.tx.capacity()), - ) - }); - } - - let mut ordered = Vec::::with_capacity(candidate_indices.len()); - for offset in 0..candidate_indices.len() { - ordered.push(candidate_indices[(start + offset) % candidate_indices.len()]); - } - ordered - }; let mut fallback_blocking_idx: Option = None; for idx in ordered_candidate_indices { @@ -804,167 +545,4 @@ impl MePool { } } - /// Send RPC_PROXY_REQ while keeping the first bound-writer path allocation-light. - /// The client byte permit follows the payload until writer completion or command drop. - pub async fn send_proxy_req_pooled( - self: &Arc, - conn_id: u64, - target_dc: i16, - client_addr: SocketAddr, - our_addr: SocketAddr, - payload: PooledBuffer, - _permit: OwnedSemaphorePermit, - proto_flags: u32, - tag_override: Option<[u8; 16]>, - ) -> Result<()> { - let tag = tag_override.or_else(|| proxy_tag_array(self.proxy_tag.as_deref())); - let Some((writer_byte_permits, writer_reserved_bytes)) = proxy_req_resident_permits( - payload.capacity(), - payload.len(), - tag.as_ref().map(|tag| tag.as_slice()), - proto_flags, - ) else { - self.stats.increment_me_writer_byte_budget_oversize_total(); - return Err(ProxyError::Proxy( - "ME writer payload residency calculation overflow".into(), - )); - }; - if writer_byte_permits as usize > self.writer_lifecycle.writer_byte_budget_permits { - self.stats.increment_me_writer_byte_budget_oversize_total(); - return Err(ProxyError::Proxy( - "ME writer payload exceeds configured byte budget".into(), - )); - } - - if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await { - let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout); - let writer_permit = match reserve_writer_bytes( - ¤t.byte_budget, - writer_byte_permits, - writer_reserved_bytes, - deadline, - &self.stats, - ) - .await - { - Ok(permit) => permit, - Err(WriterByteReserveError::TimedOut) => { - self.stats - .increment_me_writer_pick_full_total(self.writer_pick_mode()); - return Err(ProxyError::Proxy( - "ME writer byte budget full within blocking send timeout".into(), - )); - } - Err(WriterByteReserveError::Closed) => { - warn!( - writer_id = current.writer_id, - "ME writer byte budget closed" - ); - self.remove_writer_and_close_clients(current.writer_id) - .await; - return self - .send_proxy_req( - conn_id, - target_dc, - client_addr, - our_addr, - payload.as_ref(), - proto_flags, - tag.as_ref().map(|tag| tag.as_slice()), - Some(_permit), - ) - .await; - } - }; - let command = WriterCommand::ProxyReq(ProxyReqCommand { - conn_id, - client_addr, - our_addr: current_meta.our_addr, - proto_flags, - proxy_tag: tag, - payload, - _permit, - writer_permit, - }); - match current.tx.try_send(command) { - Ok(()) => { - self.note_hybrid_route_success(); - return Ok(()); - } - Err(TrySendError::Full(cmd)) => { - match reserve_writer_command_slot(¤t.tx, deadline).await { - Ok(permit) => { - permit.send(cmd); - self.note_hybrid_route_success(); - return Ok(()); - } - Err(WriterCommandReserveError::TimedOut) => { - self.stats - .increment_me_writer_pick_full_total(self.writer_pick_mode()); - return Err(ProxyError::Proxy( - "ME writer channel full within blocking send timeout".into(), - )); - } - Err(WriterCommandReserveError::Closed) => { - let Some((payload, _permit)) = proxy_req_payload_from_command(cmd) - else { - return Err(ProxyError::Proxy( - "ME writer rejected unexpected command type".into(), - )); - }; - warn!(writer_id = current.writer_id, "ME writer channel closed"); - self.remove_writer_and_close_clients(current.writer_id) - .await; - return self - .send_proxy_req( - conn_id, - target_dc, - client_addr, - our_addr, - payload.as_ref(), - proto_flags, - tag.as_ref().map(|tag| tag.as_slice()), - Some(_permit), - ) - .await; - } - } - } - Err(TrySendError::Closed(cmd)) => { - let Some((payload, _permit)) = proxy_req_payload_from_command(cmd) else { - return Err(ProxyError::Proxy( - "ME writer rejected unexpected command type".into(), - )); - }; - warn!(writer_id = current.writer_id, "ME writer channel closed"); - self.remove_writer_and_close_clients(current.writer_id) - .await; - return self - .send_proxy_req( - conn_id, - target_dc, - client_addr, - our_addr, - payload.as_ref(), - proto_flags, - tag.as_ref().map(|tag| tag.as_slice()), - Some(_permit), - ) - .await; - } - } - } - - self.send_proxy_req( - conn_id, - target_dc, - client_addr, - our_addr, - payload.as_ref(), - proto_flags, - tag.as_ref().map(|tag| tag.as_slice()), - Some(_permit), - ) - .await - } } diff --git a/src/transport/middle_proxy/send/bound.rs b/src/transport/middle_proxy/send/bound.rs new file mode 100644 index 0000000..20b27c7 --- /dev/null +++ b/src/transport/middle_proxy/send/bound.rs @@ -0,0 +1,114 @@ +use std::net::SocketAddr; +use std::sync::Arc; + +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::mpsc::error::TrySendError; +use tracing::warn; + +use super::super::MePool; +use super::super::codec::WriterCommand; +use super::super::wire::build_proxy_req_payload; +use super::reservation::{ + WriterByteReserveError, WriterCommandReserveError, payload_permit_from_data_command, + reserve_writer_bytes, reserve_writer_command_slot, writer_send_deadline, +}; +use crate::error::{ProxyError, Result}; + +pub(super) enum BoundWriterSendOutcome { + Sent, + Retry(Option), +} + +impl MePool { + pub(super) async fn try_send_bound_writer( + self: &Arc, + conn_id: u64, + client_addr: SocketAddr, + data: &[u8], + proto_flags: u32, + tag: Option<&[u8]>, + writer_byte_permits: u32, + writer_reserved_bytes: usize, + payload_permit: Option, + ) -> Result { + let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await else { + return Ok(BoundWriterSendOutcome::Retry(payload_permit)); + }; + let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout); + let writer_permit = match reserve_writer_bytes( + ¤t.byte_budget, + writer_byte_permits, + writer_reserved_bytes, + deadline, + &self.stats, + ) + .await + { + Ok(permit) => permit, + Err(WriterByteReserveError::TimedOut) => { + self.stats + .increment_me_writer_pick_full_total(self.writer_pick_mode()); + return Err(ProxyError::Proxy( + "ME writer byte budget full within blocking send timeout".into(), + )); + } + Err(WriterByteReserveError::Closed) => { + warn!( + writer_id = current.writer_id, + "ME writer byte budget closed" + ); + self.remove_writer_and_close_clients(current.writer_id) + .await; + return Ok(BoundWriterSendOutcome::Retry(payload_permit)); + } + }; + let payload = build_proxy_req_payload( + conn_id, + client_addr, + current_meta.our_addr, + data, + tag, + proto_flags, + ); + let command = WriterCommand::Data { + payload, + _permit: payload_permit, + writer_permit, + }; + match current.tx.try_send(command) { + Ok(()) => { + self.note_hybrid_route_success(); + Ok(BoundWriterSendOutcome::Sent) + } + Err(TrySendError::Full(cmd)) => { + match reserve_writer_command_slot(¤t.tx, deadline).await { + Ok(permit) => { + permit.send(cmd); + self.note_hybrid_route_success(); + return Ok(BoundWriterSendOutcome::Sent); + } + Err(WriterCommandReserveError::TimedOut) => { + self.stats + .increment_me_writer_pick_full_total(self.writer_pick_mode()); + return Err(ProxyError::Proxy( + "ME writer channel full within blocking send timeout".into(), + )); + } + Err(WriterCommandReserveError::Closed) => {} + } + let payload_permit = payload_permit_from_data_command(cmd); + warn!(writer_id = current.writer_id, "ME writer channel closed"); + self.remove_writer_and_close_clients(current.writer_id) + .await; + Ok(BoundWriterSendOutcome::Retry(payload_permit)) + } + Err(TrySendError::Closed(cmd)) => { + let payload_permit = payload_permit_from_data_command(cmd); + warn!(writer_id = current.writer_id, "ME writer channel closed"); + self.remove_writer_and_close_clients(current.writer_id) + .await; + Ok(BoundWriterSendOutcome::Retry(payload_permit)) + } + } + } +} diff --git a/src/transport/middle_proxy/send/pooled.rs b/src/transport/middle_proxy/send/pooled.rs new file mode 100644 index 0000000..13cc4e9 --- /dev/null +++ b/src/transport/middle_proxy/send/pooled.rs @@ -0,0 +1,182 @@ +use std::net::SocketAddr; +use std::sync::Arc; + +use tokio::sync::OwnedSemaphorePermit; +use tokio::sync::mpsc::error::TrySendError; +use tracing::warn; + +use super::super::MePool; +use super::super::codec::{ProxyReqCommand, WriterCommand}; +use super::reservation::{ + WriterByteReserveError, WriterCommandReserveError, proxy_req_payload_from_command, + proxy_req_resident_permits, proxy_tag_array, reserve_writer_bytes, + reserve_writer_command_slot, writer_send_deadline, +}; +use crate::error::{ProxyError, Result}; +use crate::stream::PooledBuffer; + +impl MePool { + /// Send RPC_PROXY_REQ while keeping the first bound-writer path allocation-light. + /// The client byte permit follows the payload until writer completion or command drop. + pub async fn send_proxy_req_pooled( + self: &Arc, + conn_id: u64, + target_dc: i16, + client_addr: SocketAddr, + our_addr: SocketAddr, + payload: PooledBuffer, + _permit: OwnedSemaphorePermit, + proto_flags: u32, + tag_override: Option<[u8; 16]>, + ) -> Result<()> { + let tag = tag_override.or_else(|| proxy_tag_array(self.proxy_tag.as_deref())); + let Some((writer_byte_permits, writer_reserved_bytes)) = proxy_req_resident_permits( + payload.capacity(), + payload.len(), + tag.as_ref().map(|tag| tag.as_slice()), + proto_flags, + ) else { + self.stats.increment_me_writer_byte_budget_oversize_total(); + return Err(ProxyError::Proxy( + "ME writer payload residency calculation overflow".into(), + )); + }; + if writer_byte_permits as usize > self.writer_lifecycle.writer_byte_budget_permits { + self.stats.increment_me_writer_byte_budget_oversize_total(); + return Err(ProxyError::Proxy( + "ME writer payload exceeds configured byte budget".into(), + )); + } + + if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await { + let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout); + let writer_permit = match reserve_writer_bytes( + ¤t.byte_budget, + writer_byte_permits, + writer_reserved_bytes, + deadline, + &self.stats, + ) + .await + { + Ok(permit) => permit, + Err(WriterByteReserveError::TimedOut) => { + self.stats + .increment_me_writer_pick_full_total(self.writer_pick_mode()); + return Err(ProxyError::Proxy( + "ME writer byte budget full within blocking send timeout".into(), + )); + } + Err(WriterByteReserveError::Closed) => { + warn!( + writer_id = current.writer_id, + "ME writer byte budget closed" + ); + self.remove_writer_and_close_clients(current.writer_id) + .await; + return self + .send_proxy_req( + conn_id, + target_dc, + client_addr, + our_addr, + payload.as_ref(), + proto_flags, + tag.as_ref().map(|tag| tag.as_slice()), + Some(_permit), + ) + .await; + } + }; + let command = WriterCommand::ProxyReq(ProxyReqCommand { + conn_id, + client_addr, + our_addr: current_meta.our_addr, + proto_flags, + proxy_tag: tag, + payload, + _permit, + writer_permit, + }); + match current.tx.try_send(command) { + Ok(()) => { + self.note_hybrid_route_success(); + return Ok(()); + } + Err(TrySendError::Full(cmd)) => { + match reserve_writer_command_slot(¤t.tx, deadline).await { + Ok(permit) => { + permit.send(cmd); + self.note_hybrid_route_success(); + return Ok(()); + } + Err(WriterCommandReserveError::TimedOut) => { + self.stats + .increment_me_writer_pick_full_total(self.writer_pick_mode()); + return Err(ProxyError::Proxy( + "ME writer channel full within blocking send timeout".into(), + )); + } + Err(WriterCommandReserveError::Closed) => { + let Some((payload, _permit)) = proxy_req_payload_from_command(cmd) + else { + return Err(ProxyError::Proxy( + "ME writer rejected unexpected command type".into(), + )); + }; + warn!(writer_id = current.writer_id, "ME writer channel closed"); + self.remove_writer_and_close_clients(current.writer_id) + .await; + return self + .send_proxy_req( + conn_id, + target_dc, + client_addr, + our_addr, + payload.as_ref(), + proto_flags, + tag.as_ref().map(|tag| tag.as_slice()), + Some(_permit), + ) + .await; + } + } + } + Err(TrySendError::Closed(cmd)) => { + let Some((payload, _permit)) = proxy_req_payload_from_command(cmd) else { + return Err(ProxyError::Proxy( + "ME writer rejected unexpected command type".into(), + )); + }; + warn!(writer_id = current.writer_id, "ME writer channel closed"); + self.remove_writer_and_close_clients(current.writer_id) + .await; + return self + .send_proxy_req( + conn_id, + target_dc, + client_addr, + our_addr, + payload.as_ref(), + proto_flags, + tag.as_ref().map(|tag| tag.as_slice()), + Some(_permit), + ) + .await; + } + } + } + + self.send_proxy_req( + conn_id, + target_dc, + client_addr, + our_addr, + payload.as_ref(), + proto_flags, + tag.as_ref().map(|tag| tag.as_slice()), + Some(_permit), + ) + .await + } +} diff --git a/src/transport/middle_proxy/send/reservation.rs b/src/transport/middle_proxy/send/reservation.rs new file mode 100644 index 0000000..951e505 --- /dev/null +++ b/src/transport/middle_proxy/send/reservation.rs @@ -0,0 +1,143 @@ +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError, mpsc}; + +use super::super::codec::{WriterBytePermit, WriterCommand}; +use crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES; +use crate::stats::Stats; +use crate::stream::PooledBuffer; + +const RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES: usize = 27; +pub(super) const LEGACY_PROXY_REQ_SOURCE_CAPACITY_OVERHEAD_BYTES: usize = 128; + +pub(super) enum WriterCommandReserveError { + Closed, + TimedOut, +} + +pub(super) enum WriterByteReserveError { + Closed, + TimedOut, +} + +pub(super) fn proxy_tag_array(tag: Option<&[u8]>) -> Option<[u8; 16]> { + tag.and_then(|tag| <[u8; 16]>::try_from(tag).ok()) +} + +pub(super) fn proxy_req_payload_from_command( + cmd: WriterCommand, +) -> Option<(PooledBuffer, OwnedSemaphorePermit)> { + match cmd { + WriterCommand::ProxyReq(command) => Some((command.payload, command._permit)), + _ => None, + } +} + +pub(super) fn payload_permit_from_data_command( + cmd: WriterCommand, +) -> Option { + match cmd { + WriterCommand::Data { _permit, .. } => _permit, + _ => None, + } +} + +pub(super) async fn reserve_writer_command_slot( + tx: &mpsc::Sender, + deadline: Option, +) -> std::result::Result, WriterCommandReserveError> { + let reserve = tx.clone().reserve_owned(); + match deadline { + Some(deadline) => { + match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), reserve) + .await + { + Ok(Ok(permit)) => Ok(permit), + Ok(Err(_)) => Err(WriterCommandReserveError::Closed), + Err(_) => Err(WriterCommandReserveError::TimedOut), + } + } + None => reserve.await.map_err(|_| WriterCommandReserveError::Closed), + } +} + +pub(super) fn writer_send_deadline(wait: Option) -> Option { + wait.map(|wait| Instant::now() + wait) +} + +fn writer_resident_permits( + source_capacity: usize, + encoded_payload_len: usize, +) -> Option<(u32, usize)> { + let resident_bytes = source_capacity + .checked_add(encoded_payload_len)? + .checked_add(RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES)?; + let permits = resident_bytes.div_ceil(ME_WRITER_BYTE_PERMIT_UNIT_BYTES); + let permits = u32::try_from(permits).ok()?; + let reserved_bytes = (permits as usize).checked_mul(ME_WRITER_BYTE_PERMIT_UNIT_BYTES)?; + Some(( + permits.max(1), + reserved_bytes.max(ME_WRITER_BYTE_PERMIT_UNIT_BYTES), + )) +} + +pub(super) fn proxy_req_resident_permits( + source_capacity: usize, + data_len: usize, + proxy_tag: Option<&[u8]>, + proto_flags: u32, +) -> Option<(u32, usize)> { + writer_resident_permits( + source_capacity, + super::super::wire::proxy_req_payload_len(data_len, proxy_tag, proto_flags), + ) +} + +pub(super) fn try_reserve_writer_bytes( + byte_budget: &Arc, + permits: u32, + reserved_bytes: usize, + stats: &Arc, +) -> std::result::Result { + byte_budget + .clone() + .try_acquire_many_owned(permits) + .map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone())) +} + +pub(super) async fn reserve_writer_bytes( + byte_budget: &Arc, + permits: u32, + reserved_bytes: usize, + deadline: Option, + stats: &Arc, +) -> std::result::Result { + match try_reserve_writer_bytes(byte_budget, permits, reserved_bytes, stats) { + Ok(permit) => return Ok(permit), + Err(TryAcquireError::Closed) => return Err(WriterByteReserveError::Closed), + Err(TryAcquireError::NoPermits) => { + stats.increment_me_writer_byte_budget_wait_total(); + } + } + + let acquire = byte_budget.clone().acquire_many_owned(permits); + match deadline { + Some(deadline) => { + match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), acquire) + .await + { + Ok(Ok(permit)) => Ok(WriterBytePermit::new(permit, reserved_bytes, stats.clone())), + Ok(Err(_)) => Err(WriterByteReserveError::Closed), + Err(_) => { + stats.increment_me_writer_byte_budget_timeout_total(); + Err(WriterByteReserveError::TimedOut) + } + } + } + None => acquire + .await + .map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone())) + .map_err(|_| WriterByteReserveError::Closed), + } +} diff --git a/src/transport/middle_proxy/send/selection.rs b/src/transport/middle_proxy/send/selection.rs index 5c75b13..fd85491 100644 --- a/src/transport/middle_proxy/send/selection.rs +++ b/src/transport/middle_proxy/send/selection.rs @@ -1,3 +1,4 @@ +use std::cmp::Reverse; use std::collections::{HashMap, HashSet}; use std::sync::atomic::Ordering; @@ -7,6 +8,7 @@ use super::{ IDLE_WRITER_PENALTY_HIGH_SECS, IDLE_WRITER_PENALTY_MID_SECS, PICK_PENALTY_DEGRADED, PICK_PENALTY_DRAINING, PICK_PENALTY_STALE, PICK_PENALTY_WARM, }; +use crate::config::MeWriterPickMode; impl MePool { pub(super) async fn candidate_indices_for_dc( @@ -184,4 +186,94 @@ impl MePool { } ordered } + + pub(super) async fn ordered_candidate_indices( + &self, + mut candidate_indices: Vec, + writers_snapshot: &[super::super::pool::MeWriter], + pick_mode: MeWriterPickMode, + ) -> Vec { + let pick_sample_size = self.writer_pick_sample_size(); + let writer_ids: Vec = candidate_indices + .iter() + .map(|idx| writers_snapshot[*idx].id) + .collect(); + let writer_idle_since = self + .registry + .writer_idle_since_for_writer_ids(&writer_ids) + .await; + let now_epoch_secs = Self::now_epoch_secs(); + let start = self.rr.fetch_add(1, Ordering::Relaxed) as usize % candidate_indices.len(); + if pick_mode == MeWriterPickMode::P2c { + return self.p2c_ordered_candidate_indices( + &candidate_indices, + writers_snapshot, + &writer_idle_since, + now_epoch_secs, + start, + pick_sample_size, + ); + } + + if self + .writer_selection_policy + .me_deterministic_writer_sort + .load(Ordering::Relaxed) + { + candidate_indices.sort_by(|lhs, rhs| { + let left = &writers_snapshot[*lhs]; + let right = &writers_snapshot[*rhs]; + let left_key = ( + self.writer_contour_rank_for_selection(left), + (left.generation < self.current_generation()) as usize, + left.degraded.load(Ordering::Relaxed) as usize, + self.writer_idle_rank_for_selection( + left, + &writer_idle_since, + now_epoch_secs, + ), + Reverse(left.tx.capacity()), + left.addr, + left.id, + ); + let right_key = ( + self.writer_contour_rank_for_selection(right), + (right.generation < self.current_generation()) as usize, + right.degraded.load(Ordering::Relaxed) as usize, + self.writer_idle_rank_for_selection( + right, + &writer_idle_since, + now_epoch_secs, + ), + Reverse(right.tx.capacity()), + right.addr, + right.id, + ); + left_key.cmp(&right_key) + }); + } else { + candidate_indices.sort_by_key(|idx| { + let writer = &writers_snapshot[*idx]; + let degraded = writer.degraded.load(Ordering::Relaxed); + let stale = (writer.generation < self.current_generation()) as usize; + ( + self.writer_contour_rank_for_selection(writer), + stale, + degraded as usize, + self.writer_idle_rank_for_selection( + writer, + &writer_idle_since, + now_epoch_secs, + ), + Reverse(writer.tx.capacity()), + ) + }); + } + + let mut ordered = Vec::::with_capacity(candidate_indices.len()); + for offset in 0..candidate_indices.len() { + ordered.push(candidate_indices[(start + offset) % candidate_indices.len()]); + } + ordered + } } diff --git a/src/transport/middle_proxy/tests/pool_refill_security_tests.rs b/src/transport/middle_proxy/tests/pool_refill_security_tests.rs index db9245d..80f859c 100644 --- a/src/transport/middle_proxy/tests/pool_refill_security_tests.rs +++ b/src/transport/middle_proxy/tests/pool_refill_security_tests.rs @@ -1,15 +1,19 @@ use std::collections::HashMap; -use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::sync::Arc; use std::sync::atomic::Ordering; use std::time::{Duration, Instant}; use crate::config::{GeneralConfig, MeRouteNoWriterMode, MeSocksKdfPolicy, MeWriterPickMode}; use crate::crypto::SecureRandom; +use crate::network::IpFamily; use crate::network::probe::NetworkDecision; use crate::stats::Stats; -use super::pool::MePool; +use super::pool::{ + MePool, ReinitStatusSnapshot, WriterContour, WriterRole, +}; +use super::pool_writer_security_tests::make_pool_with_decision; async fn make_pool() -> Arc { let general = GeneralConfig::default(); @@ -177,6 +181,91 @@ async fn refill_does_not_queue_a_removed_dc_target() { assert_eq!(pool.refill_pending.load(Ordering::Acquire), 0); } +#[tokio::test(flavor = "current_thread")] +async fn refill_accepts_enabled_nonpreferred_family_without_multipath() { + let pool = make_pool_with_decision(NetworkDecision { + ipv4_me: true, + ipv6_me: true, + effective_prefer: 4, + effective_multipath: false, + ..NetworkDecision::default() + }) + .await; + let v4_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let v6_addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443); + pool.update_proxy_maps( + HashMap::from([(2, vec![(v4_addr.ip(), v4_addr.port())])]), + Some(HashMap::from([( + 2, + vec![(v6_addr.ip(), v6_addr.port())], + )])), + ) + .await; + + pool.trigger_immediate_refill_for_dc(v6_addr, 2); + + assert_eq!(pool.refill_states.lock().len(), 1); + assert_eq!(pool.refill_running.load(Ordering::Acquire), 1); + pool.begin_shutdown(); + tokio::time::timeout(Duration::from_secs(1), async { + while pool.refill_running.load(Ordering::Acquire) != 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert!(pool.refill_states.lock().is_empty()); +} + +#[tokio::test(flavor = "current_thread")] +async fn stale_endpoint_revision_cancels_queued_warm_refill_before_connect() { + let pool = make_pool().await; + let old_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 30)), 443); + let new_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 31)), 443); + pool.update_proxy_maps( + HashMap::from([(2, vec![(old_addr.ip(), old_addr.port())])]), + None, + ) + .await; + let old_revision = pool.endpoint_snapshot.load().revision; + let pending_generation = pool.current_generation().saturating_add(1); + pool.reinit.status.store(Arc::new(ReinitStatusSnapshot { + active_generation: pool.current_generation(), + warm_generations: vec![pending_generation], + pending_hardswap_generation: pending_generation, + pending_hardswap_started_at_epoch_secs: 1, + pending_hardswap_map_hash: 1, + pending_hardswap_endpoint_revision: old_revision, + inflight: 1, + })); + pool.trigger_immediate_refill_for_role( + old_addr, + WriterRole { + dc: 2, + family: IpFamily::V4, + generation: pending_generation, + contour: WriterContour::Warm, + }, + ); + assert_eq!(pool.refill_states.lock().len(), 1); + + pool.update_proxy_maps( + HashMap::from([(2, vec![(new_addr.ip(), new_addr.port())])]), + None, + ) + .await; + tokio::time::timeout(Duration::from_secs(1), async { + while pool.refill_running.load(Ordering::Acquire) != 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + assert!(pool.refill_states.lock().is_empty()); + assert_eq!(pool.stats.get_me_reconnect_attempts(), 0); +} + #[tokio::test(flavor = "current_thread")] async fn refill_preserves_bounded_pending_cardinality_and_cleans_up_before_first_poll() { let pool = make_pool().await; diff --git a/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs index 84def94..2315196 100644 --- a/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs +++ b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs @@ -1,4 +1,4 @@ -use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering}; use std::time::{Duration, Instant}; @@ -8,8 +8,9 @@ use tokio_util::sync::CancellationToken; use super::codec::WriterCommand; use super::pool::{MeWriter, WriterContour, WriterOpenIntent}; -use super::pool_writer_security_tests::make_pool; +use super::pool_writer_security_tests::{make_pool, make_pool_with_decision}; use super::registry::ConnMeta; +use crate::network::probe::NetworkDecision; fn unregistered_writer( pool: &Arc, @@ -128,6 +129,225 @@ async fn normal_active_publication_cannot_race_past_the_family_floor() { ); } +#[tokio::test] +async fn stale_same_family_writers_do_not_satisfy_current_endpoint_coverage() { + let pool = make_pool().await; + let current_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let stale_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)), 443); + pool.update_proxy_maps( + std::collections::HashMap::from([( + 2, + vec![(current_addr.ip(), current_addr.port())], + )]), + None, + ) + .await; + pool.floor_runtime + .me_adaptive_floor_cpu_cores_override + .store(1, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .store(1, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_global + .store(1, Ordering::Relaxed); + let generation = pool.current_generation(); + let required = pool.required_writers_for_dc_with_floor_mode(1, false); + pool.writers + .write() + .await + .extend((1..=required.saturating_mul(2)).map(|writer_id| { + unregistered_writer( + &pool, + writer_id as u64, + stale_addr, + generation, + WriterContour::Active, + ) + })); + + assert!( + pool.can_open_writer_for_contour( + WriterContour::Active, + WriterOpenIntent::Coverage, + 2, + current_addr, + ) + .await + ); + assert!( + !pool + .can_open_writer_for_contour( + WriterContour::Active, + WriterOpenIntent::Coverage, + 2, + stale_addr, + ) + .await + ); +} + +#[tokio::test] +async fn covered_group_does_not_consume_another_group_coverage_slot() { + let pool = make_pool().await; + let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; + pool.floor_runtime + .me_adaptive_floor_cpu_cores_override + .store(1, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .store(1, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_global + .store(1, Ordering::Relaxed); + let generation = pool.current_generation(); + let required = pool.required_writers_for_dc_with_floor_mode(1, false); + pool.writers + .write() + .await + .extend((1..=required).map(|writer_id| { + unregistered_writer( + &pool, + writer_id as u64, + addr, + generation, + WriterContour::Active, + ) + })); + + assert!( + !pool + .can_open_writer_for_contour( + WriterContour::Active, + WriterOpenIntent::Coverage, + 2, + addr, + ) + .await + ); +} + +#[tokio::test] +async fn normal_active_publication_allows_adaptive_growth_above_family_floor() { + let pool = make_pool().await; + let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; + let generation = pool.current_generation(); + let writers = (1..=3) + .map(|writer_id| { + unregistered_writer( + &pool, + writer_id, + addr, + generation, + WriterContour::Active, + ) + }) + .collect::>(); + let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active); + + assert!( + pool.authorize_writer_publication_capacity( + &candidate, + WriterContour::Active, + WriterOpenIntent::Normal, + &writers, + ) + .is_ok() + ); +} + +#[tokio::test] +async fn normal_active_publication_rejects_configured_contour_cap() { + let pool = make_pool().await; + let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; + pool.floor_runtime + .me_adaptive_floor_cpu_cores_override + .store(1, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_per_core + .store(3, Ordering::Relaxed); + pool.floor_runtime + .me_adaptive_floor_max_active_writers_global + .store(3, Ordering::Relaxed); + let generation = pool.current_generation(); + let writers = (1..=3) + .map(|writer_id| { + unregistered_writer( + &pool, + writer_id, + addr, + generation, + WriterContour::Active, + ) + }) + .collect::>(); + let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active); + + assert!( + pool.authorize_writer_publication_capacity( + &candidate, + WriterContour::Active, + WriterOpenIntent::Normal, + &writers, + ) + .is_err() + ); +} + +#[tokio::test] +async fn nonpreferred_enabled_family_retains_writer_publication_authority() { + let pool = make_pool_with_decision(NetworkDecision { + ipv4_me: true, + ipv6_me: true, + effective_prefer: 4, + effective_multipath: false, + ..NetworkDecision::default() + }) + .await; + let v4_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let v6_addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(v4_addr.ip(), v4_addr.port())])]), + Some(std::collections::HashMap::from([( + 2, + vec![(v6_addr.ip(), v6_addr.port())], + )])), + ) + .await; + let candidate = unregistered_writer( + &pool, + 1, + v6_addr, + pool.current_generation(), + WriterContour::Active, + ); + + assert!( + pool.authorize_writer_publication_capacity( + &candidate, + WriterContour::Active, + WriterOpenIntent::Coverage, + &[], + ) + .is_ok() + ); +} + #[tokio::test] async fn successful_writer_publication_is_fully_visible_and_removable() { let pool = make_pool().await; diff --git a/src/transport/middle_proxy/tests/pool_writer_security_tests.rs b/src/transport/middle_proxy/tests/pool_writer_security_tests.rs index e47524e..5b16b79 100644 --- a/src/transport/middle_proxy/tests/pool_writer_security_tests.rs +++ b/src/transport/middle_proxy/tests/pool_writer_security_tests.rs @@ -17,6 +17,15 @@ use crate::stats::Stats; /// Builds an isolated ME pool for writer-state tests. pub(super) async fn make_pool() -> Arc { + make_pool_with_decision(NetworkDecision { + ipv4_me: true, + ..NetworkDecision::default() + }) + .await +} + +/// Builds an isolated ME pool with an explicit network-family policy. +pub(super) async fn make_pool_with_decision(decision: NetworkDecision) -> Arc { let general = GeneralConfig::default(); MePool::new( @@ -35,10 +44,7 @@ pub(super) async fn make_pool() -> Arc { HashMap::new(), HashMap::new(), None, - NetworkDecision { - ipv4_me: true, - ..NetworkDecision::default() - }, + decision, None, Arc::new(SecureRandom::new()), Arc::new(Stats::new()), diff --git a/src/web/manager.rs b/src/web/manager.rs index 1a0e03b..d7d21dd 100644 --- a/src/web/manager.rs +++ b/src/web/manager.rs @@ -14,7 +14,6 @@ use crate::config::{WebCarrier, WebLimitsConfig}; use crate::maestro::generation::RuntimeGeneration; use crate::web::telemetry::{WebRejectionReason, WebTelemetry}; use crate::web::trace::WebTraceStore; - // Credential maps, quotas, and token-bucket helpers remain private to the manager. mod state; // Carrier attempt metadata remains explicit and independent from HTTP parsing. diff --git a/src/web/manager/session_creation/replacement.rs b/src/web/manager/session_creation/replacement.rs index 865974b..a4e4cc8 100644 --- a/src/web/manager/session_creation/replacement.rs +++ b/src/web/manager/session_creation/replacement.rs @@ -57,6 +57,8 @@ impl WebProcessRuntime { return Err(ManagerError::Closed); }; let Some(user_registration) = user_publication.take_registration() else { + // Rollback callbacks can retire active user owners and must not run under authority. + drop(user_publication); drop(state); self.cancel_replacement(bootstrap_hash, &replacement.old_session); return Err(ManagerError::Closed); @@ -65,6 +67,8 @@ impl WebProcessRuntime { self.record_limit_hit(); self.telemetry .record_rejection(crate::web::telemetry::WebRejectionReason::SessionCapacity); + drop(user_registration); + drop(user_publication); drop(state); self.cancel_replacement(bootstrap_hash, &replacement.old_session); return Err(ManagerError::Limit); @@ -99,6 +103,7 @@ impl WebProcessRuntime { Some(user_registration), ); let Some(supersede) = replacement.old_session.prepare_carrier_supersede() else { + drop(user_publication); drop(state); self.cancel_replacement(bootstrap_hash, &replacement.old_session); session.close(crate::web::session::SessionCloseReason::Protocol);