diff --git a/src/payment/forwarding_store.rs b/src/payment/forwarding_store.rs index 645406297..86b7c814b 100644 --- a/src/payment/forwarding_store.rs +++ b/src/payment/forwarding_store.rs @@ -73,9 +73,13 @@ pub(crate) struct ForwardRecord<'a> { #[derive(Clone, Debug, PartialEq, Eq)] struct ForwardedPaymentReplayMarker { id: ForwardedPaymentId, + forwarded_at_timestamp: u64, } -impl_writeable_tlv_based!(ForwardedPaymentReplayMarker, { (0, id, required) }); +impl_writeable_tlv_based!(ForwardedPaymentReplayMarker, { + (0, id, required), + (2, forwarded_at_timestamp, required), +}); impl StorableObject for ForwardedPaymentReplayMarker { type Id = ForwardedPaymentId; @@ -308,15 +312,14 @@ impl ForwardingStore { })?; } - // Keep this marker after the event is handled. LDK can replay an older event after later - // events have replaced the directional retry tokens, and it provides no callback after its - // handled-event state is durable. - self.replay_markers.insert(ForwardedPaymentReplayMarker { id: forward_id }).await.map_err( - |e| { + // Must outlive the event's replay window; removed once aggregated or swept by age. + self.replay_markers + .insert(ForwardedPaymentReplayMarker { id: forward_id, forwarded_at_timestamp }) + .await + .map_err(|e| { log_error!(self.logger, "Failed to store forwarded payment replay marker: {e}"); e - }, - )?; + })?; Ok(()) } @@ -366,12 +369,30 @@ impl ForwardingStore { self.channel_pair_stats.list_page(page_token).await } - pub(crate) async fn aggregate_expired(&self, retention_secs: u64) -> Result<(u64, u64), Error> { + pub(crate) async fn aggregate_expired( + &self, retention_secs: u64, remove_markers: bool, + ) -> Result<(u64, u64), Error> { aggregate_expired_forwarded_payments( &self.details, &self.replay_markers, &self.channel_pair_stats, retention_secs, + remove_markers, + &self.logger, + ) + .await + } + + async fn prune_stale_replay_markers(&self) -> Result { + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::from_secs(0)) + .as_secs(); + prune_expired_replay_markers( + &self.details, + &self.replay_markers, + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS, + now, &self.logger, ) .await @@ -523,9 +544,9 @@ fn seconds_until_next_forwarding_aggregation(now_timestamp: u64, bucket_size_sec } async fn aggregate_forwarded_payments_and_log( - forwarding_store: &ForwardingStore, retention_secs: u64, + forwarding_store: &ForwardingStore, retention_secs: u64, remove_markers: bool, ) { - match forwarding_store.aggregate_expired(retention_secs).await { + match forwarding_store.aggregate_expired(retention_secs, remove_markers).await { Ok((pair_count, payment_count)) if pair_count > 0 => { log_debug!( forwarding_store.logger, @@ -548,39 +569,16 @@ async fn aggregate_forwarded_payments_and_log( } } +/// Marker removal only starts after one bucket width of uptime, so a replay of an event recorded +/// before a crash still finds its marker. pub(crate) async fn run_forwarded_payment_aggregation( mut stop_receiver: tokio::sync::watch::Receiver<()>, forwarding_store: Arc, retention_secs: u64, ) { - if retention_secs == 0 { - match forwarding_store.details.is_empty().await { - Ok(true) => return, - Ok(false) => {}, - Err(e) => { - log_error!( - forwarding_store.logger, - "Failed to check forwarded payment store: {}", - e - ); - }, - } - } + let started_at = tokio::time::Instant::now(); + let prune_delay = Duration::from_secs(FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS); - aggregate_forwarded_payments_and_log(&forwarding_store, retention_secs).await; - - if retention_secs == 0 { - match forwarding_store.details.is_empty().await { - Ok(true) => return, - Ok(false) => {}, - Err(e) => { - log_error!( - forwarding_store.logger, - "Failed to check forwarded payment store: {}", - e - ); - }, - } - } + aggregate_forwarded_payments_and_log(&forwarding_store, retention_secs, false).await; let period = Duration::from_secs(FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS); let now = @@ -596,26 +594,43 @@ pub(crate) async fn run_forwarded_payment_aggregation( tokio::select! { _ = stop_receiver.changed() => break, _ = interval.tick() => { - aggregate_forwarded_payments_and_log(&forwarding_store, retention_secs).await; - if retention_secs == 0 { - match forwarding_store.details.is_empty().await { - Ok(true) => break, - Ok(false) => {}, - Err(e) => { log_error!(forwarding_store.logger, "Failed to check forwarded payment store: {}", e); }, - } + let remove_markers = started_at.elapsed() >= prune_delay; + aggregate_forwarded_payments_and_log(&forwarding_store, retention_secs, remove_markers).await; + if remove_markers { + prune_stale_replay_markers_and_log(&forwarding_store).await; } } } } } +async fn prune_stale_replay_markers_and_log(forwarding_store: &ForwardingStore) { + match forwarding_store.prune_stale_replay_markers().await { + Ok(removed) if removed > 0 => { + log_debug!( + forwarding_store.logger, + "Reclaimed {} expired forwarded payment replay markers", + removed + ); + }, + Err(e) => { + log_error!( + forwarding_store.logger, + "Failed to reclaim forwarded payment replay markers: {}", + e + ); + }, + _ => {}, + } +} + /// Aggregate forwarded payments older than the configured retention period into fixed-width /// channel-pair statistics buckets. async fn aggregate_expired_forwarded_payments( forwarded_payment_store: &ForwardedPaymentStore, replay_marker_store: &ForwardedPaymentReplayMarkerStore, channel_pair_stats_store: &ChannelPairForwardingStatsStore, retention_secs: u64, - logger: &Arc, + remove_markers: bool, logger: &Arc, ) -> Result<(u64, u64), Error> { let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or(Duration::from_secs(0)).as_secs(); @@ -626,16 +641,62 @@ async fn aggregate_expired_forwarded_payments( FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS, retention_secs, now, + remove_markers, logger, ) .await } +/// Removes markers older than one bucket width whose detail, if any, is gone. Age alone isn't +/// enough: pruning a sibling's marker would leave its detail stranded in a deferred bucket. +async fn prune_expired_replay_markers( + forwarded_payment_store: &ForwardedPaymentStore, + replay_marker_store: &ForwardedPaymentReplayMarkerStore, bucket_size_secs: u64, now: u64, + logger: &Arc, +) -> Result { + if bucket_size_secs == 0 { + return Ok(0); + } + let oldest_retained_bucket_start = + (now.saturating_sub(bucket_size_secs) / bucket_size_secs).saturating_mul(bucket_size_secs); + + // `Stats` mode writes no details, so skip the per-marker lookup. + let no_details_exist = forwarded_payment_store.is_empty().await?; + + let mut expired_ids = Vec::new(); + let mut page_token = None; + loop { + let page = replay_marker_store.list_page(page_token).await?; + for marker in page.objects { + if marker.forwarded_at_timestamp >= oldest_retained_bucket_start { + continue; + } + if !no_details_exist && forwarded_payment_store.contains_key(&marker.id).await? { + continue; + } + expired_ids.push(marker.id); + } + let Some(next_page_token) = page.next_page_token else { break }; + page_token = Some(next_page_token); + } + + let mut removed = 0u64; + for id in expired_ids { + match replay_marker_store.remove(&id).await { + Ok(()) => removed += 1, + Err(e) => { + log_error!(logger, "Failed to remove replay marker {:?}: {}", id, e); + }, + } + } + Ok(removed) +} + async fn aggregate_expired_forwarded_payments_at( forwarded_payment_store: &ForwardedPaymentStore, replay_marker_store: &ForwardedPaymentReplayMarkerStore, channel_pair_stats_store: &ChannelPairForwardingStatsStore, bucket_size_secs: u64, - retention_secs: u64, now: u64, logger: &Arc, + retention_secs: u64, now: u64, remove_markers: bool, logger: &Arc, ) -> Result<(u64, u64), Error> { if bucket_size_secs == 0 { return Ok((0, 0)); @@ -771,6 +832,12 @@ async fn aggregate_expired_forwarded_payments_at( log_error!(logger, "Failed to remove forwarded payment {:?}: {}", payment_id, e); e })?; + if remove_markers { + // On failure the marker is orphaned and the sweep reclaims it later. + if let Err(e) = replay_marker_store.remove(&payment_id).await { + log_error!(logger, "Failed to remove replay marker {:?}: {}", payment_id, e); + } + } removed_payment_count += 1; } @@ -887,8 +954,10 @@ pub fn aggregate_channel_pair_stats( mod forwarding_stats_tests { use std::str::FromStr; + use lightning::io; use lightning::util::persist::{ - KVStore, KVSTORE_NAMESPACE_KEY_ALPHABET, KVSTORE_NAMESPACE_KEY_MAX_LEN, + KVStore, PaginatedKVStore, PaginatedListResponse, KVSTORE_NAMESPACE_KEY_ALPHABET, + KVSTORE_NAMESPACE_KEY_MAX_LEN, }; use lightning::util::ser::{Readable, Writeable}; @@ -950,13 +1019,199 @@ mod forwarding_stats_tests { (forwarded_payment_store, replay_marker_store, channel_pair_stats_store, logger, kv_store) } + /// Fails `remove` for one specific key; everything else delegates to the inner store. + struct FailRemoveForKey { + inner: InMemoryStore, + failing_secondary_namespace: &'static str, + failing_key: String, + } + + impl KVStore for FailRemoveForKey { + fn read( + &self, primary_namespace: &str, secondary_namespace: &str, key: &str, + ) -> impl std::future::Future, io::Error>> + 'static + Send { + self.inner.read(primary_namespace, secondary_namespace, key) + } + fn write( + &self, primary_namespace: &str, secondary_namespace: &str, key: &str, buf: Vec, + ) -> impl std::future::Future> + 'static + Send { + self.inner.write(primary_namespace, secondary_namespace, key, buf) + } + fn remove( + &self, primary_namespace: &str, secondary_namespace: &str, key: &str, lazy: bool, + ) -> impl std::future::Future> + 'static + Send { + let fail = + secondary_namespace == self.failing_secondary_namespace && key == self.failing_key; + let fut: std::pin::Pin< + Box> + Send>, + > = if fail { + Box::pin(async { Err(io::Error::new(io::ErrorKind::Other, "remove failed")) }) + } else { + Box::pin(self.inner.remove(primary_namespace, secondary_namespace, key, lazy)) + }; + fut + } + fn list( + &self, primary_namespace: &str, secondary_namespace: &str, + ) -> impl std::future::Future, io::Error>> + 'static + Send { + self.inner.list(primary_namespace, secondary_namespace) + } + } + + impl PaginatedKVStore for FailRemoveForKey { + fn list_paginated( + &self, primary_namespace: &str, secondary_namespace: &str, + page_token: Option, + ) -> impl std::future::Future> + 'static + Send + { + self.inner.list_paginated(primary_namespace, secondary_namespace, page_token) + } + } + + #[tokio::test] + async fn sweep_continues_after_a_marker_removal_failure() { + let ids = [ForwardedPaymentId([1; 32]), ForwardedPaymentId([2; 32])]; + // Whichever marker the sweep reaches first, the other must still be reclaimed. + for failing_id in ids { + let kv_store: Arc = Arc::new(DynStoreWrapper(FailRemoveForKey { + inner: InMemoryStore::new(), + failing_secondary_namespace: "replay_markers", + failing_key: failing_id.encode_to_hex_str(), + })); + let logger = Arc::new(Logger::new_log_facade()); + let forwarded_payment_store = ForwardedPaymentStore::new( + Vec::new(), + KeepNoEntries, + "test_forwarded_payments".to_string(), + "details".to_string(), + Arc::clone(&kv_store), + Arc::clone(&logger), + ); + let replay_marker_store = ForwardedPaymentReplayMarkerStore::new( + Vec::new(), + KeepNoEntries, + "test_forwarded_payments".to_string(), + "replay_markers".to_string(), + Arc::clone(&kv_store), + Arc::clone(&logger), + ); + for id in ids { + replay_marker_store + .insert(ForwardedPaymentReplayMarker { id, forwarded_at_timestamp: 0 }) + .await + .unwrap(); + } + + let removed = prune_expired_replay_markers( + &forwarded_payment_store, + &replay_marker_store, + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS, + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS * 10, + &logger, + ) + .await + .unwrap(); + + assert_eq!(removed, 1); + assert!(replay_marker_store.contains_key(&failing_id).await.unwrap()); + let other_id = ids.into_iter().find(|id| *id != failing_id).unwrap(); + assert!(!replay_marker_store.contains_key(&other_id).await.unwrap()); + } + } + + #[tokio::test] + async fn marker_removal_failure_does_not_abort_the_rest_of_the_pass() { + let failing_payment = forwarded_payment(1, 850, 110, 100, 10); + let other_payment = forwarded_payment(2, 851, 220, 200, 20); + let failing_key = failing_payment.id().encode_to_hex_str(); + + let kv_store: Arc = Arc::new(DynStoreWrapper(FailRemoveForKey { + inner: InMemoryStore::new(), + failing_secondary_namespace: "replay_markers", + failing_key, + })); + let logger = Arc::new(Logger::new_log_facade()); + let primary_namespace = "test_forwarded_payments"; + let forwarded_payment_store = ForwardedPaymentStore::new( + Vec::new(), + KeepNoEntries, + primary_namespace.to_string(), + "details".to_string(), + Arc::clone(&kv_store), + Arc::clone(&logger), + ); + let replay_marker_store = ForwardedPaymentReplayMarkerStore::new( + Vec::new(), + KeepNoEntries, + primary_namespace.to_string(), + "replay_markers".to_string(), + Arc::clone(&kv_store), + Arc::clone(&logger), + ); + let channel_pair_stats_store = ChannelPairForwardingStatsStore::new( + Vec::new(), + KeepNoEntries, + primary_namespace.to_string(), + "pair_stats".to_string(), + Arc::clone(&kv_store), + Arc::clone(&logger), + ); + + insert_completed_payment( + &forwarded_payment_store, + &replay_marker_store, + failing_payment.clone(), + ) + .await; + insert_completed_payment( + &forwarded_payment_store, + &replay_marker_store, + other_payment.clone(), + ) + .await; + + let result = aggregate_expired_forwarded_payments_at( + &forwarded_payment_store, + &replay_marker_store, + &channel_pair_stats_store, + 60, + 60, + 1_000, + true, + &logger, + ) + .await + .unwrap(); + assert_eq!( + result, + (1, 2), + "both details are aggregated despite the marker-removal failure" + ); + + assert!(forwarded_payment_store.get(&failing_payment.id()).await.unwrap().is_none()); + assert!(forwarded_payment_store.get(&other_payment.id()).await.unwrap().is_none()); + + assert!( + replay_marker_store.contains_key(&failing_payment.id()).await.unwrap(), + "the marker whose removal failed must survive, orphaned, for the sweep to reclaim" + ); + assert!( + !replay_marker_store.contains_key(&other_payment.id()).await.unwrap(), + "the other payment's marker must still be removed normally" + ); + } + async fn insert_completed_payment( forwarded_payment_store: &TestForwardedPaymentStore, replay_marker_store: &TestReplayMarkerStore, payment: ForwardedPaymentDetails, ) { let id = payment.id(); + let forwarded_at_timestamp = payment.forwarded_at_timestamp; forwarded_payment_store.insert(payment).await.unwrap(); - replay_marker_store.insert(ForwardedPaymentReplayMarker { id }).await.unwrap(); + replay_marker_store + .insert(ForwardedPaymentReplayMarker { id, forwarded_at_timestamp }) + .await + .unwrap(); } fn forwarded_payment( @@ -1056,6 +1311,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1106,6 +1362,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1121,7 +1378,10 @@ mod forwarding_stats_tests { ); replay_marker_store - .insert(ForwardedPaymentReplayMarker { id: payment.id() }) + .insert(ForwardedPaymentReplayMarker { + id: payment.id(), + forwarded_at_timestamp: payment.forwarded_at_timestamp, + }) .await .unwrap(); assert_eq!( @@ -1132,6 +1392,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1160,19 +1421,31 @@ mod forwarding_stats_tests { forwarding_store.details.insert(payment.clone()).await.unwrap(); forwarding_store .replay_markers - .insert(ForwardedPaymentReplayMarker { id: payment.id() }) + .insert(ForwardedPaymentReplayMarker { + id: payment.id(), + forwarded_at_timestamp: payment.forwarded_at_timestamp, + }) .await .unwrap(); - let (_stop_sender, stop_receiver) = tokio::sync::watch::channel(()); + let (stop_sender, stop_receiver) = tokio::sync::watch::channel(()); - tokio::time::timeout( - Duration::from_secs(1), - run_forwarded_payment_aggregation(stop_receiver, Arc::clone(&forwarding_store), 0), - ) + let handle = tokio::spawn(run_forwarded_payment_aggregation( + stop_receiver, + Arc::clone(&forwarding_store), + 0, + )); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + if forwarding_store.details.is_empty().await.unwrap() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) .await .unwrap(); - assert!(forwarding_store.details.is_empty().await.unwrap()); + assert!(forwarding_store.replay_markers.contains_key(&payment.id()).await.unwrap()); let bucket_id = channel_pair_stats_id(&payment.prev_channel_id, &payment.next_channel_id, 0); assert_eq!( @@ -1185,6 +1458,140 @@ mod forwarding_stats_tests { .payment_count, 1 ); + assert!(!handle.is_finished(), "the loop must keep running once drained"); + + stop_sender.send(()).unwrap(); + tokio::time::timeout(Duration::from_secs(1), handle).await.unwrap().unwrap(); + } + + #[tokio::test] + async fn loop_stays_alive_in_stats_mode_when_stores_start_empty() { + let kv_store: Arc = Arc::new(DynStoreWrapper(InMemoryStore::new())); + let logger = Arc::new(Logger::new_log_facade()); + let forwarding_store = Arc::new(ForwardingStore::new( + Vec::new(), + ForwardedPaymentTrackingMode::Stats, + kv_store, + logger, + )); + let (stop_sender, stop_receiver) = tokio::sync::watch::channel(()); + + let handle = tokio::spawn(run_forwarded_payment_aggregation( + stop_receiver, + Arc::clone(&forwarding_store), + 0, + )); + tokio::time::sleep(Duration::from_millis(200)).await; + assert!( + !handle.is_finished(), + "loop exited on startup; markers written later will never be reclaimed" + ); + + stop_sender.send(()).unwrap(); + tokio::time::timeout(Duration::from_secs(1), handle).await.unwrap().unwrap(); + } + + #[tokio::test] + async fn startup_pass_does_not_sweep_an_old_marker() { + let kv_store: Arc = Arc::new(DynStoreWrapper(InMemoryStore::new())); + let logger = Arc::new(Logger::new_log_facade()); + let forwarding_store = Arc::new(ForwardingStore::new( + Vec::new(), + ForwardedPaymentTrackingMode::Stats, + kv_store, + Arc::clone(&logger), + )); + + let marker_id = ForwardedPaymentId([7; 32]); + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::from_secs(0)) + .as_secs(); + forwarding_store + .replay_markers + .insert(ForwardedPaymentReplayMarker { + id: marker_id, + forwarded_at_timestamp: now + .saturating_sub(FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS * 2), + }) + .await + .unwrap(); + + let (stop_sender, stop_receiver) = tokio::sync::watch::channel(()); + let handle = tokio::spawn(run_forwarded_payment_aggregation( + stop_receiver, + Arc::clone(&forwarding_store), + 0, + )); + + tokio::time::sleep(Duration::from_millis(200)).await; + assert!( + forwarding_store.replay_markers.contains_key(&marker_id).await.unwrap(), + "the startup pass must not sweep a marker before the uptime delay elapses" + ); + + stop_sender.send(()).unwrap(); + tokio::time::timeout(Duration::from_secs(1), handle).await.unwrap().unwrap(); + } + + /// The startup pass must not remove a marker, even when it aggregates the paired detail. + #[tokio::test] + async fn startup_pass_does_not_remove_a_paired_marker_in_detailed_mode() { + let kv_store: Arc = Arc::new(DynStoreWrapper(InMemoryStore::new())); + let logger = Arc::new(Logger::new_log_facade()); + let forwarding_store = Arc::new(ForwardingStore::new( + Vec::new(), + ForwardedPaymentTrackingMode::Detailed, + kv_store, + Arc::clone(&logger), + )); + + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::from_secs(0)) + .as_secs(); + let payment = forwarded_payment( + 1, + now.saturating_sub(FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS * 3), + 110, + 100, + 10, + ); + forwarding_store.details.insert(payment.clone()).await.unwrap(); + forwarding_store + .replay_markers + .insert(ForwardedPaymentReplayMarker { + id: payment.id(), + forwarded_at_timestamp: payment.forwarded_at_timestamp, + }) + .await + .unwrap(); + + let (stop_sender, stop_receiver) = tokio::sync::watch::channel(()); + let handle = tokio::spawn(run_forwarded_payment_aggregation( + stop_receiver, + Arc::clone(&forwarding_store), + 0, + )); + + tokio::time::timeout(Duration::from_secs(1), async { + loop { + if forwarding_store.details.is_empty().await.unwrap() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + assert!( + forwarding_store.replay_markers.contains_key(&payment.id()).await.unwrap(), + "the startup pass aggregated the detail but must leave its marker for the delayed sweep" + ); + + stop_sender.send(()).unwrap(); + tokio::time::timeout(Duration::from_secs(1), handle).await.unwrap().unwrap(); } #[tokio::test] @@ -1382,6 +1789,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1485,6 +1893,7 @@ mod forwarding_stats_tests { 60, 60, 900, + true, &logger, ) .await, @@ -1502,6 +1911,216 @@ mod forwarding_stats_tests { ); } + #[tokio::test] + async fn old_replay_markers_are_reclaimed_regardless_of_tracking_mode() { + for mode in [ForwardedPaymentTrackingMode::Stats, ForwardedPaymentTrackingMode::Detailed] { + let kv_store: Arc = Arc::new(DynStoreWrapper(InMemoryStore::new())); + let logger = Arc::new(Logger::new_log_facade()); + let forwarding_store = + ForwardingStore::new(Vec::new(), mode.clone(), Arc::clone(&kv_store), logger); + + let prev_channel_id = ChannelId([1; 32]); + let next_channel_id = ChannelId([2; 32]); + const FORWARDS: u64 = 50; + for htlc_id in 0..FORWARDS { + let prev_htlcs = [InboundHTLCLocator { + channel_id: prev_channel_id, + htlc_id: Some(htlc_id), + amount_msat: Some(1_000), + user_channel_id: None, + node_id: None, + }]; + let next_htlcs = [OutboundHTLCLocator { + channel_id: next_channel_id, + amount_msat: Some(999), + user_channel_id: None, + node_id: None, + }]; + forwarding_store + .record_forward(ForwardRecord { + prev_htlcs: &prev_htlcs, + next_htlcs: &next_htlcs, + total_fee_earned_msat: Some(1), + skimmed_fee_msat: None, + claim_from_onchain_tx: false, + outbound_amount_forwarded_msat: 999, + }) + .await + .unwrap(); + } + assert!( + !forwarding_store.replay_markers.is_empty().await.unwrap(), + "{mode:?}: forwarding should have written markers" + ); + + let real_now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or(Duration::from_secs(0)) + .as_secs(); + let far_future = real_now + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS * 3; + let (_, aggregated_removed) = aggregate_expired_forwarded_payments_at( + &forwarding_store.details, + &forwarding_store.replay_markers, + &forwarding_store.channel_pair_stats, + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS, + 0, + far_future, + true, + &forwarding_store.logger, + ) + .await + .unwrap(); + let swept = prune_expired_replay_markers( + &forwarding_store.details, + &forwarding_store.replay_markers, + FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS, + far_future, + &forwarding_store.logger, + ) + .await + .unwrap(); + + assert_eq!( + aggregated_removed + swept, + FORWARDS, + "{mode:?}: every expired marker should be reclaimed, one way or the other" + ); + assert!( + forwarding_store.replay_markers.is_empty().await.unwrap(), + "{mode:?}: no replay markers should remain" + ); + } + } + + /// Markers from the current or previous bucket must survive; older ones are reclaimed. + #[tokio::test] + async fn prune_expired_replay_markers_respects_the_bucket_cutoff() { + let (forwarded_payment_store, replay_marker_store, _, logger, _) = test_stores_with_kv(); + let bucket_size_secs = FORWARDED_PAYMENT_AGGREGATION_BUCKET_SIZE_SECS; + let now = bucket_size_secs * 10; + + let recent_id = ForwardedPaymentId([1; 32]); + let old_id = ForwardedPaymentId([2; 32]); + replay_marker_store + .insert(ForwardedPaymentReplayMarker { + id: recent_id, + forwarded_at_timestamp: now - bucket_size_secs / 2, + }) + .await + .unwrap(); + replay_marker_store + .insert(ForwardedPaymentReplayMarker { + id: old_id, + forwarded_at_timestamp: now - bucket_size_secs * 2, + }) + .await + .unwrap(); + + let removed = prune_expired_replay_markers( + &forwarded_payment_store, + &replay_marker_store, + bucket_size_secs, + now, + &logger, + ) + .await + .unwrap(); + + assert_eq!(removed, 1); + assert!(replay_marker_store.contains_key(&recent_id).await.unwrap()); + assert!(!replay_marker_store.contains_key(&old_id).await.unwrap()); + } + + /// Pruning by age alone would strand a sibling detail in a deferred bucket. + #[tokio::test] + async fn pruning_never_orphans_a_sibling_detail_in_a_deferred_bucket() { + let (forwarded_payment_store, replay_marker_store, channel_pair_stats_store, logger) = + test_stores(); + + // Same bucket: `straggler` crashed mid-write (no marker), `sibling` is complete. + let straggler = forwarded_payment(1, 850, 110, 100, 10); + let sibling = forwarded_payment(2, 851, 220, 200, 20); + forwarded_payment_store.insert(straggler.clone()).await.unwrap(); + insert_completed_payment(&forwarded_payment_store, &replay_marker_store, sibling.clone()) + .await; + + let now = 10_000; + let bucket_size_secs = 60; + let retention_secs = 60; + + let aggregate_result = aggregate_expired_forwarded_payments_at( + &forwarded_payment_store, + &replay_marker_store, + &channel_pair_stats_store, + bucket_size_secs, + retention_secs, + now, + true, + &logger, + ) + .await + .unwrap(); + assert_eq!(aggregate_result, (0, 0), "the whole bucket must be deferred"); + + let removed = prune_expired_replay_markers( + &forwarded_payment_store, + &replay_marker_store, + bucket_size_secs, + now, + &logger, + ) + .await + .unwrap(); + assert_eq!(removed, 0, "the sibling's marker still guards a live detail"); + + assert_eq!( + forwarded_payment_store.get(&sibling.id()).await.unwrap(), + Some(sibling.clone()) + ); + assert_eq!( + forwarded_payment_store.get(&straggler.id()).await.unwrap(), + Some(straggler.clone()) + ); + assert!(replay_marker_store.contains_key(&sibling.id()).await.unwrap()); + + replay_marker_store + .insert(ForwardedPaymentReplayMarker { + id: straggler.id(), + forwarded_at_timestamp: straggler.forwarded_at_timestamp, + }) + .await + .unwrap(); + + let aggregate_result = aggregate_expired_forwarded_payments_at( + &forwarded_payment_store, + &replay_marker_store, + &channel_pair_stats_store, + bucket_size_secs, + retention_secs, + now, + true, + &logger, + ) + .await + .unwrap(); + assert_eq!(aggregate_result, (1, 2), "the now-complete bucket aggregates normally"); + assert!(forwarded_payment_store.get(&sibling.id()).await.unwrap().is_none()); + assert!(forwarded_payment_store.get(&straggler.id()).await.unwrap().is_none()); + assert!(!replay_marker_store.contains_key(&sibling.id()).await.unwrap()); + assert!(!replay_marker_store.contains_key(&straggler.id()).await.unwrap()); + + let removed = prune_expired_replay_markers( + &forwarded_payment_store, + &replay_marker_store, + bucket_size_secs, + now, + &logger, + ) + .await + .unwrap(); + assert_eq!(removed, 0, "nothing left for the sweep once aggregation already removed both"); + } + #[tokio::test] async fn zero_retention_cleans_up_after_the_current_bucket_closes() { let (forwarded_payment_store, replay_marker_store, channel_pair_stats_store, logger) = @@ -1518,6 +2137,7 @@ mod forwarding_stats_tests { 60, 0, 899, + true, &logger, ) .await, @@ -1536,6 +2156,7 @@ mod forwarding_stats_tests { 60, 0, 900, + true, &logger, ) .await, @@ -1568,6 +2189,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1613,6 +2235,7 @@ mod forwarding_stats_tests { 60, 60, 1_000, + true, &logger, ) .await, @@ -1673,6 +2296,7 @@ mod forwarding_stats_tests { 3_600, 7_200, 15_000, + true, &logger, ) .await,