diff --git a/README.md b/README.md index 65fc59b5..a7355e53 100644 --- a/README.md +++ b/README.md @@ -425,6 +425,8 @@ Included: count and value limitations set on channel creation. * Liquidity checks: HTLCs are only forwarded if the node has sufficient liquidity in the mocked channel. +* Retries: a failed payment is re-routed up to 5 times, avoiding the + channels that rejected it, before it is reported as failed. Not included: * Channel reserve: the required minimum reserve balance is not diff --git a/simln-lib/src/lib.rs b/simln-lib/src/lib.rs index 05e7e1ac..4d6102cb 100755 --- a/simln-lib/src/lib.rs +++ b/simln-lib/src/lib.rs @@ -466,7 +466,9 @@ pub enum PaymentOutcome { NotDispatched, /// The payment was dispatched but its final status could not be determined. TrackPaymentFailed, - /// The payment failed at the provided index in the path. + /// The payment failed at the provided index in the path. Simulated nodes only report this for payments sent to + /// a specific route, which are not retried; payments that they route themselves report + /// [`PaymentOutcome::RetriesExhausted`] or [`PaymentOutcome::RouteNotFound`] instead. IndexFailure(usize), } diff --git a/simln-lib/src/sim_node.rs b/simln-lib/src/sim_node.rs index 931f937f..834aa1c6 100755 --- a/simln-lib/src/sim_node.rs +++ b/simln-lib/src/sim_node.rs @@ -492,6 +492,15 @@ impl SimulatedChannel { } } +/// The outcome of a payment dispatched with [`SimNetwork::dispatch_payment`], reported to the sending node. +#[derive(Debug)] +pub struct HtlcOutcome { + pub result: PaymentResult, + /// The short channel id of the channel that rejected the htlc, set when the payment was failed by the network + /// while the htlc was being added along its route. + pub failed_channel: Option, +} + /// SimNetwork represents a high level network coordinator that is responsible for the task of actually propagating /// payments through the simulated network. #[async_trait] @@ -504,7 +513,7 @@ pub trait SimNetwork: Send + Sync { route: Route, custom_records: Option, payment_hash: PaymentHash, - sender: Sender>, + sender: Sender>, ); /// Looks up a node in the simulated network and a list of its channel capacities. @@ -516,13 +525,21 @@ pub trait SimNetwork: Send + Sync { type LdkNetworkGraph = NetworkGraph>; type Scorer = ProbabilisticScorer, Arc>; +/// The maximum number of times that a failed payment is re-routed before it is failed. +const MAX_RETRIES: u32 = 5; + struct InFlightPayment { /// The channel used to report payment results to. - track_payment_receiver: Receiver>, + track_payment_receiver: Receiver>, /// The path the payment was dispatched on. /// This should be set to `None` if no payment path was found and the payment /// was not dispatched. path: Option, + /// The amount that the payment is for, re-routed to `dest` when an attempt fails. + amount_msat: u64, + /// The destination that a failed payment is re-routed to. `None` for payments sent to a specific route, which + /// are not retried. + dest: Option, } /// A wrapper struct used to implement the LightningNode trait (can be thought of as "the" lightning node). Passes @@ -605,7 +622,9 @@ impl SimNode { Entry::Vacant(vacant) => vacant.insert(InFlightPayment { track_payment_receiver: receiver, path: Some(route.paths[0].clone()), // TODO: MPP payments? we check in dispatch_payment - // should probably only pass a single path to dispatch + // should probably only pass a single path to dispatch + amount_msat: route.get_total_amount(), + dest: None, }), }; @@ -619,6 +638,161 @@ impl SimNode { Ok(()) } + + /// Drives a dispatched payment to its outcome, re-routing it over paths that avoid the channels that have + /// rejected it until it succeeds, no route is left or [`MAX_RETRIES`] retries have been made. Each attempt is + /// reported to the scorer. + async fn resolve_payment( + &self, + payment_hash: PaymentHash, + in_flight: InFlightPayment, + ) -> Result { + let InFlightPayment { + mut track_payment_receiver, + mut path, + amount_msat, + dest, + } = in_flight; + + let mut failed_channels = vec![]; + let mut htlc_count = 0; + let mut retries = 0; + + loop { + let outcome = track_payment_receiver.await.map_err(|e| { + LightningError::TrackPaymentError(format!("channel receive err: {}", e)) + })??; + htlc_count += outcome.result.htlc_count; + + let attempt = match path { + Some(ref attempt) => attempt, + // No route was found for the payment, so nothing was dispatched to score or retry. + None => { + if outcome.result.payment_outcome != PaymentOutcome::RouteNotFound { + return Err(LightningError::TrackPaymentError( + "payment outcome was not RouteNotFound, but no path was provided" + .to_string(), + )); + } + + return Ok(outcome.result); + }, + }; + + let duration = self.clock.now().duration_since(UNIX_EPOCH).map_err(|e| { + log::error!("Failed to get duration: {}", e); + LightningError::SystemTimeConversionError(e) + })?; + if outcome.result.payment_outcome == PaymentOutcome::Success { + self.scorer + .write() + .map_err(|e| { + LightningError::TrackPaymentError(format!("scorer lock poisoned: {e}")) + })? + .payment_path_successful(attempt, duration); + + return Ok(PaymentResult { + htlc_count, + payment_outcome: PaymentOutcome::Success, + }); + } + + // If the network did not blame a channel for the failure, there is nothing to learn from it and + // nothing for a retry to avoid. + let failed_channel = match outcome.failed_channel { + Some(failed_channel) => failed_channel, + None => { + return Ok(PaymentResult { + htlc_count, + ..outcome.result + }) + }, + }; + + self.scorer + .write() + .map_err(|e| { + LightningError::TrackPaymentError(format!("scorer lock poisoned: {e}")) + })? + .payment_path_failed(attempt, failed_channel, duration); + failed_channels.push(failed_channel); + + // Payments sent to a specific route are not retried, because the route was not ours to choose. + let dest = match dest { + Some(dest) => dest, + None => { + return Ok(PaymentResult { + htlc_count, + ..outcome.result + }) + }, + }; + + if retries == MAX_RETRIES { + return Ok(PaymentResult { + htlc_count, + payment_outcome: PaymentOutcome::RetriesExhausted, + }); + } + retries += 1; + + log::debug!( + "Retrying payment {} ({retries}/{MAX_RETRIES}), avoiding channels: {failed_channels:?}.", + hex::encode(payment_hash.0), + ); + + // Re-route on the blocking pool for the same reason as the payment's first route: pathfinding is + // pure CPU work that would otherwise hold up the scheduler thread. + let route = { + let source = self.info.pubkey; + let pathfinding_graph = self.pathfinding_graph.clone(); + let scorer = self.scorer.clone(); + let previously_failed_channels = failed_channels.clone(); + + tokio::task::spawn_blocking(move || -> Result<_, LightningError> { + let scorer = scorer.read().map_err(|e| { + LightningError::TrackPaymentError(format!("scorer lock poisoned: {e}")) + })?; + Ok(find_payment_route( + &source, + dest, + amount_msat, + &pathfinding_graph, + &scorer, + previously_failed_channels, + )) + }) + .await + .map_err(|e| { + LightningError::TrackPaymentError(format!("pathfinding task failed: {e}")) + })?? + }; + + let route = match route { + Ok(route) => route, + Err(e) => { + log::trace!("Could not find route to retry payment: {:?}.", e); + + return Ok(PaymentResult { + htlc_count, + payment_outcome: PaymentOutcome::RouteNotFound, + }); + }, + }; + + let (sender, receiver) = channel(); + track_payment_receiver = receiver; + path = Some(route.paths[0].clone()); + + self.network.lock().await.dispatch_payment( + self.info.pubkey, + route, + None, // Default custom records. + payment_hash, + sender, + ); + } + } } /// Produces the node info for a mocked node, filling in the features that the simulator requires. @@ -635,22 +809,26 @@ fn node_info(pubkey: PublicKey, alias: String) -> NodeInfo { } /// Uses LDK's pathfinding algorithm with default parameters to find a path from source to destination, with no -/// restrictions on fee budget. +/// restrictions on fee budget. The path will avoid the channels in `previously_failed_channels`. fn find_payment_route( source: &PublicKey, dest: PublicKey, amount_msat: u64, pathfinding_graph: &LdkNetworkGraph, scorer: &Scorer, + previously_failed_channels: Vec, ) -> Result { + let mut payment_params = PaymentParameters::from_node_id(dest, 0) + // TODO: set non-zero value to support MPP. + .with_max_path_count(1) + // Allow sending htlcs up to 50% of the channel's capacity. + .with_max_channel_saturation_power_of_half(1); + payment_params.previously_failed_channels = previously_failed_channels; + find_route( source, &RouteParameters { - payment_params: PaymentParameters::from_node_id(dest, 0) - // TODO: set non-zero value to support MPP. - .with_max_path_count(1) - // Allow sending htlcs up to 50% of the channel's capacity. - .with_max_channel_saturation_power_of_half(1), + payment_params, final_value_msat: amount_msat, max_total_routing_fee_msat: None, }, @@ -708,6 +886,7 @@ impl LightningNode for SimNode { amount_msat, &pathfinding_graph, &scorer, + Vec::new(), )) }) .await @@ -733,9 +912,12 @@ impl LightningNode for SimNode { Err(e) => { log::trace!("Could not find path for payment: {:?}.", e); - if let Err(e) = sender.send(Ok(PaymentResult { - htlc_count: 0, - payment_outcome: PaymentOutcome::RouteNotFound, + if let Err(e) = sender.send(Ok(HtlcOutcome { + result: PaymentResult { + htlc_count: 0, + payment_outcome: PaymentOutcome::RouteNotFound, + }, + failed_channel: None, })) { log::error!("Could not send payment result: {:?}.", e); } @@ -743,7 +925,9 @@ impl LightningNode for SimNode { entry.insert(InFlightPayment { track_payment_receiver: receiver, path: None, // TODO: how to handle non-MPP support (where would we do - // paths in the world where we have them?). + // paths in the world where we have them?). + amount_msat, + dest: None, }); return Ok(payment_hash); @@ -759,7 +943,9 @@ impl LightningNode for SimNode { entry.insert(InFlightPayment { track_payment_receiver: receiver, path: Some(route.paths[0].clone()), // TODO: how to handle non-MPP support (where would we do - // paths in the world where we have them?). + // paths in the world where we have them?). + amount_msat, + dest: Some(dest), }); // If we did successfully obtain a route, dispatch the payment through the network and then report success. @@ -782,53 +968,22 @@ impl LightningNode for SimNode { hash: &PaymentHash, listener: Listener, ) -> Result { - match self.in_flight.lock().await.remove(hash) { - Some(in_flight) => { - select! { - biased; - _ = listener => Err( - LightningError::TrackPaymentError("shutdown during payment tracking".to_string()), - ), - - // If we get a payment result back, remove from our in flight set of payments and return the result. - res = in_flight.track_payment_receiver => { - let track_result = res.map_err(|e| LightningError::TrackPaymentError(format!("channel receive err: {}", e)))?; - if let Ok(ref payment_result) = track_result { - let duration = match self.clock.now().duration_since(UNIX_EPOCH) { - Ok(d) => d, - Err(e) => { - log::error!("Failed to get duration: {}", e); - return Err(LightningError::SystemTimeConversionError(e)); - } - }; - match &in_flight.path { - Some(path) => { - let mut scorer = self.scorer.write().map_err(|e| { - LightningError::TrackPaymentError(format!("scorer lock poisoned: {e}")) - })?; - if payment_result.payment_outcome == PaymentOutcome::Success { - scorer.payment_path_successful(path, duration); - } else if let PaymentOutcome::IndexFailure(index) = payment_result.payment_outcome { - scorer.payment_path_failed(path, index as u64, duration); - } - }, - None => { - if payment_result.payment_outcome != PaymentOutcome::RouteNotFound { - return Err(LightningError::TrackPaymentError( - "payment outcome was not RouteNotFound, but no path was provided".to_string(), - ))?; - } - } - } - } - track_result - }, - } + let in_flight = match self.in_flight.lock().await.remove(hash) { + Some(in_flight) => in_flight, + None => { + return Err(LightningError::TrackPaymentError(format!( + "payment hash {} not found", + hex::encode(hash.0), + ))) }, - None => Err(LightningError::TrackPaymentError(format!( - "payment hash {} not found", - hex::encode(hash.0), - ))), + }; + + select! { + biased; + _ = listener => Err( + LightningError::TrackPaymentError("shutdown during payment tracking".to_string()), + ), + res = self.resolve_payment(*hash, in_flight) => res, } } @@ -1243,7 +1398,7 @@ impl SimNetwork for SimGraph { route: Route, custom_records: Option, payment_hash: PaymentHash, - sender: Sender>, + sender: Sender>, ) { // Expect only one path (right now), with the intention to support multiple in future. if route.paths.len() != 1 { @@ -1256,9 +1411,12 @@ impl SimNetwork for SimGraph { None => { log::warn!("Find route did not return expected number of paths."); - if let Err(e) = sender.send(Ok(PaymentResult { - htlc_count: 0, - payment_outcome: PaymentOutcome::RouteNotFound, + if let Err(e) = sender.send(Ok(HtlcOutcome { + result: PaymentResult { + htlc_count: 0, + payment_outcome: PaymentOutcome::RouteNotFound, + }, + failed_channel: None, })) { log::error!("Could not send payment result: {:?}.", e); } @@ -1312,6 +1470,16 @@ impl SimNetwork for SimGraph { } } +/// A htlc that was rejected while being added along its route. +struct HtlcFailure { + /// The index of the last hop that the htlc was added on, which must be failed back, or `None` if the very first + /// add was rejected. + fail_idx: Option, + /// The index of the hop whose channel rejected the htlc. + failed_hop: usize, + error: ForwardingError, +} + /// Adds htlcs to the simulation state along the path provided. If it encounters a critical error, /// it will be returned in the outer Result to signal that the simulation should shut down. If a /// forwarding error happens, it will return the index in the path from which to fail back htlcs @@ -1340,7 +1508,7 @@ async fn add_htlcs( interceptors: Vec>, custom_records: CustomRecords, shutdown_listener: Listener, -) -> Result, ForwardingError)>, CriticalError> { +) -> Result, CriticalError> { let mut outgoing_node = source; let mut outgoing_amount = route.fee_msat() + route.final_value_msat(); let mut outgoing_cltv = route.hops.iter().map(|hop| hop.cltv_expiry_delta).sum(); @@ -1375,9 +1543,15 @@ async fn add_htlcs( }, )? { Ok(idx) => idx, - // If we couldn't add to this HTLC, we only need to fail back from the preceding hop, so we don't - // have to progress our fail_idx. - Err(e) => return Ok(Err((fail_idx, e))), + // This hop's channel rejected the HTLC, so it is to blame for the failure. We only need to fail + // back from the preceding hop, so we don't have to progress our fail_idx. + Err(e) => { + return Ok(Err(HtlcFailure { + fail_idx, + failed_hop: i, + error: e, + })) + }, }; // If the HTLC was successfully added, then we'll need to remove the HTLC from this channel if we fail, @@ -1401,9 +1575,14 @@ async fn add_htlcs( outgoing_amount - hop.fee_msat, hop.fee_msat, )? { - // If we haven't met forwarding conditions for the next channel's policy, then we fail at - // the current index, because we've already added the HTLC as outgoing. - return Ok(Err((fail_idx, e))); + // If we haven't met forwarding conditions for the next channel's policy, that channel is + // to blame for the failure. We fail at the current index, because we've already added the + // HTLC as outgoing. + return Ok(Err(HtlcFailure { + fail_idx, + failed_hop: i + 1, + error: e, + })); } } } @@ -1447,7 +1626,15 @@ async fn add_htlcs( // Collect any custom records (if any) set by the interceptor(s) for the outgoing link. let attached_custom_records = match intercepted_res { Ok(records) => records, - Err(fwd_err) => return Ok(Err((fail_idx, fwd_err))), + // The interceptor acts for the node at this hop, which is deciding whether to forward the HTLC on the + // next channel (or to accept it, if it is the receiver), so that channel is to blame for the failure. + Err(fwd_err) => { + return Ok(Err(HtlcFailure { + fail_idx, + failed_hop: (i + 1).min(last_hop), + error: fwd_err, + })) + }, }; // Once we've taken the "hop" to the destination pubkey, it becomes the source of the next outgoing htlc and @@ -1533,7 +1720,7 @@ struct PropagatePaymentRequest { source: PublicKey, route: Path, payment_hash: PaymentHash, - sender: Sender>, + sender: Sender>, interceptors: Vec>, custom_records: CustomRecords, shutdown_signal: (Trigger, Listener), @@ -1571,12 +1758,21 @@ async fn propagate_payment(request: PropagatePaymentRequest) { request.shutdown_signal.0.trigger(); log::error!("Could not remove htlcs from channel: {e}."); } - PaymentResult { - htlc_count: 1, - payment_outcome: PaymentOutcome::Success, + HtlcOutcome { + result: PaymentResult { + htlc_count: 1, + payment_outcome: PaymentOutcome::Success, + }, + failed_channel: None, } }, - Ok(Err((fail_idx, fwd_err))) => { + Ok(Err(HtlcFailure { + fail_idx, + failed_hop, + error: fwd_err, + })) => { + // Grab the channel to blame before the route is consumed by failing the htlc back. + let failed_channel = request.route.hops[failed_hop].short_channel_id; // If we partially added HTLCs along the route, we need to fail them back to the source to clean up our partial // state. It's possible that we failed with the very first add, and then we don't need to clean anything up. if let Some(resolution_idx) = fail_idx { @@ -1597,12 +1793,15 @@ async fn propagate_payment(request: PropagatePaymentRequest) { } log::debug!( - "Forwarding failure for simulated payment {}: {fwd_err}", + "Forwarding failure for simulated payment {} at hop {failed_hop}: {fwd_err}", hex::encode(request.payment_hash.0) ); - PaymentResult { - htlc_count: 0, - payment_outcome: PaymentOutcome::IndexFailure(fail_idx.unwrap_or(0)), + HtlcOutcome { + result: PaymentResult { + htlc_count: 1, + payment_outcome: PaymentOutcome::IndexFailure(fail_idx.unwrap_or(0)), + }, + failed_channel: Some(failed_channel), } }, Err(critical_err) => { @@ -1611,9 +1810,12 @@ async fn propagate_payment(request: PropagatePaymentRequest) { "Critical error in simulated payment {}: {critical_err}", hex::encode(request.payment_hash.0) ); - PaymentResult { - htlc_count: 0, - payment_outcome: PaymentOutcome::Unknown, + HtlcOutcome { + result: PaymentResult { + htlc_count: 0, + payment_outcome: PaymentOutcome::Unknown, + }, + failed_channel: None, } }, }; @@ -2242,7 +2444,7 @@ mod tests { route: Route, custom_records: Option, payment_hash: PaymentHash, - sender: Sender>, + sender: Sender>, ); async fn lookup_node(&self, node: &PublicKey) -> Result<(NodeInfo, Vec), LightningError>; @@ -2303,7 +2505,7 @@ mod tests { route: Route, _: Option, _, - sender: Sender>| { + sender: Sender>| { // If we've reached dispatch, we must have at least one path, grab the last hop to match the // receiver. let receiver = route.paths[0].hops.last().unwrap().pubkey; @@ -2321,7 +2523,12 @@ mod tests { panic!("unknown mocked receiver"); }; - sender.send(Ok(result)).unwrap(); + sender + .send(Ok(HtlcOutcome { + result, + failed_channel: None, + })) + .unwrap(); }, ); @@ -2521,19 +2728,36 @@ mod tests { source: PublicKey, dest: PublicKey, amt: u64, - ) -> (Route, Result) { - let route = - find_payment_route(&source, dest, amt, &self.routing_graph, &self.scorer).unwrap(); + ) -> (Route, Result) { + let route = find_payment_route( + &source, + dest, + amt, + &self.routing_graph, + &self.scorer, + Vec::new(), + ) + .unwrap(); + + let outcome = self.dispatch_route(source, route.clone()).await; + (route, outcome) + } + // Dispatches the route provided through the network and returns the outcome reported for its htlc. + async fn dispatch_route( + &mut self, + source: PublicKey, + route: Route, + ) -> Result { let (sender, receiver) = oneshot::channel(); self.graph - .dispatch_payment(source, route.clone(), None, PaymentHash([1; 32]), sender); + .dispatch_payment(source, route, None, PaymentHash([1; 32]), sender); let payment_result = timeout(Duration::from_millis(10), receiver).await; // Assert that we receive from the channel or fail. assert!(payment_result.is_ok()); - (route, payment_result.unwrap().unwrap()) + payment_result.unwrap().unwrap() } // Sets the balance on the channel to the tuple provided, used to arrange liquidity for testing. @@ -2767,6 +2991,239 @@ mod tests { assert!(matches!(result.payment_outcome, PaymentOutcome::Success)); } + /// Tests that a failed htlc blames the channel that rejected it, for each place that a htlc can be rejected: + /// the receiving node's interceptor, the next channel's forwarding policy and the channel refusing the add. + #[tokio::test] + async fn test_failure_reports_blamed_channel() { + let chan_capacity = 500_000_000; + + // Interceptor that fails htlcs at the receiving node only. + let mut interceptor = MockTestInterceptor::new(); + interceptor.expect_intercept_htlc().returning(|req| { + if req.outgoing_channel_id.is_none() { + Ok(Err(ForwardingError::InterceptorError( + "receiver rejected".into(), + ))) + } else { + Ok(Ok(CustomRecords::default())) + } + }); + interceptor.expect_notify_resolution().returning(|_| Ok(())); + + let interceptor: Arc = Arc::new(interceptor); + let mut test_kit = + DispatchPaymentTestKit::new(chan_capacity, vec![interceptor], CustomRecords::default()) + .await; + let initial_balances = test_kit.channel_balances().await; + let (alice, dave) = (test_kit.nodes[0], test_kit.nodes[3]); + + // Rejected by the receiver, which blames the channel that the receiver was reached on. + let (route, outcome) = test_kit.send_test_payment(alice, dave, 20_000).await; + let scids: Vec = route.paths[0] + .hops + .iter() + .map(|hop| hop.short_channel_id) + .collect(); + assert_eq!(outcome.unwrap().failed_channel, Some(scids[2])); + assert_eq!(test_kit.channel_balances().await, initial_balances); + + // Rejected by Carol --> Dave's forwarding policy, because Bob forwards with an insufficient cltv delta, + // after htlcs were added on Alice --> Bob and Bob --> Carol. + let mut short_cltv = route.clone(); + short_cltv.paths[0].hops[1].cltv_expiry_delta = 39; + let outcome = test_kit.dispatch_route(alice, short_cltv).await; + assert_eq!(outcome.unwrap().failed_channel, Some(scids[2])); + assert_eq!(test_kit.channel_balances().await, initial_balances); + + // Rejected by Bob --> Carol, which has no liquidity to forward the htlc. + test_kit + .set_channel_balance(&ShortChannelID::from(1), (0, chan_capacity)) + .await; + let outcome = test_kit.dispatch_route(alice, route).await; + assert_eq!(outcome.unwrap().failed_channel, Some(scids[1])); + + test_kit.shutdown.0.trigger(); + test_kit.graph.tasks.close(); + test_kit.graph.tasks.wait().await; + } + + /// Creates `routes` routes from Alice to Dave, each through its own intermediate node, with all of the + /// capacity of each channel on the side of the first node. With two routes (via Bob and via Carol) the network + /// is: + /// + /// Alice --(0)-- Bob --(1)-- Dave + /// | | + /// +---(2)-- Carol --(3)----+ + /// + /// Each route charges a higher base fee than the one before it, so that pathfinding has a strict preference + /// between them and picks them in order rather than breaking a tie. The nodes are returned in the order Alice, + /// the intermediate nodes, Dave. + fn create_parallel_channels( + routes: usize, + capacity_msat: u64, + ) -> (Vec, Vec) { + let nodes: Vec = (0..routes + 2).map(|_| get_random_keypair().1).collect(); + let (alice, dave) = (nodes[0], nodes[routes + 1]); + + let policy = |pubkey: PublicKey, base_fee: u64| ChannelPolicy { + pubkey, + alias: String::default(), + max_htlc_count: 483, + max_in_flight_msat: capacity_msat, + min_htlc_size_msat: 1, + max_htlc_size_msat: capacity_msat, + cltv_expiry_delta: 40, + base_fee, + fee_rate_prop: 1000, + }; + + let channels = nodes[1..=routes] + .iter() + .enumerate() + .flat_map(|(route, mid)| { + let base_fee = 1000 * (route as u64 + 1); + [(alice, *mid, base_fee), (*mid, dave, base_fee)] + }) + .enumerate() + .map(|(i, (node_1, node_2, base_fee))| SimulatedChannel { + capacity_msat, + short_channel_id: ShortChannelID::from(i as u64), + node_1: ChannelState::new(policy(node_1, base_fee), capacity_msat), + node_2: ChannelState::new(policy(node_2, base_fee), 0), + exclude_capacity: false, + }) + .collect(); + + (channels, nodes) + } + + /// Creates a simulated node for Alice in a network of `routes` routes to Dave (see + /// [`create_parallel_channels`]), returning the node and the network's pubkeys. + fn create_parallel_node( + routes: usize, + capacity: u64, + interceptors: Vec>, + ) -> (SimNode, Vec) { + let (channels, nodes) = create_parallel_channels(routes, capacity); + let clock = Arc::new(SimulationClock::new(SystemTime::now())); + let routing_graph = + Arc::new(populate_network_graph(channels.clone(), clock.clone()).unwrap()); + + let graph = SimGraph::new( + channels, + TaskTracker::new(), + interceptors, + CustomRecords::default(), + triggered::trigger(), + ) + .expect("could not create test graph"); + + let node = SimNode::new( + node_info(nodes[0], String::default()), + Arc::new(Mutex::new(graph)), + routing_graph, + clock, + ) + .unwrap(); + + (node, nodes) + } + + /// Drains the channel provided so that its first node cannot forward any htlcs. + async fn drain_channel(node: &SimNode, scid: u64) { + let network = node.network.lock().await; + let mut channels = network.channels.lock().await; + let channel = channels.get_mut(&ShortChannelID::from(scid)).unwrap(); + channel.node_2.local_balance_msat += channel.node_1.local_balance_msat; + channel.node_1.local_balance_msat = 0; + } + + /// Sends a payment from Alice to the destination provided and waits for its outcome. + async fn send_and_track( + node: &SimNode, + dest: PublicKey, + amount_msat: u64, + ) -> PaymentResult { + let payment_hash = node.send_payment(dest, amount_msat).await.unwrap(); + let (_, shutdown_listener) = triggered::trigger(); + node.track_payment(&payment_hash, shutdown_listener) + .await + .unwrap() + } + + /// Tests that a payment rejected by a channel is retried over a route that avoids it. + #[tokio::test] + async fn test_retry_avoids_failed_channel() { + let capacity = 1_000_000; + let (node, nodes) = create_parallel_node(2, capacity, vec![]); + let dave = nodes[3]; + + // Drain Bob --> Dave, so that the cheapest route (and so the one that pathfinding picks first) cannot + // carry the payment and only a retry through Carol can succeed. + drain_channel(&node, 1).await; + + let result = send_and_track(&node, dave, 100_000).await; + assert!( + matches!(result.payment_outcome, PaymentOutcome::Success), + "unexpected outcome: {:?}", + result.payment_outcome + ); + + // One htlc was dispatched for the attempt through Bob and one for the retry through Carol. + assert_eq!(result.htlc_count, 2); + } + + /// Tests that a payment with no route left to retry over is failed as soon as pathfinding gives up. + #[tokio::test] + async fn test_retry_without_route_fails_payment() { + let capacity = 1_000_000; + let (node, nodes) = create_parallel_node(1, capacity, vec![]); + let dave = nodes[2]; + + // Drain the only channel that reaches Dave, so that the payment fails and its retry has no route left. + drain_channel(&node, 1).await; + + let result = send_and_track(&node, dave, 100_000).await; + assert!( + matches!(result.payment_outcome, PaymentOutcome::RouteNotFound), + "unexpected outcome: {:?}", + result.payment_outcome + ); + assert_eq!(result.htlc_count, 1); + } + + /// Tests that a payment that is rejected on every attempt is failed once its retry budget is spent. + #[tokio::test] + async fn test_retries_exhausted() { + let capacity = 1_000_000; + + // Interceptor that rejects htlcs at the receiving node, so that every attempt fails. + let mut interceptor = MockTestInterceptor::new(); + interceptor.expect_intercept_htlc().returning(|req| { + if req.outgoing_channel_id.is_none() { + Ok(Err(ForwardingError::InterceptorError( + "receiver rejected".into(), + ))) + } else { + Ok(Ok(CustomRecords::default())) + } + }); + interceptor.expect_notify_resolution().returning(|_| Ok(())); + + // Give the network more routes than the retry budget, so that pathfinding always has one left. + let routes = MAX_RETRIES as usize + 1; + let (node, nodes) = create_parallel_node(routes, capacity, vec![Arc::new(interceptor)]); + let dave = nodes[routes + 1]; + + let result = send_and_track(&node, dave, 100_000).await; + assert!( + matches!(result.payment_outcome, PaymentOutcome::RetriesExhausted), + "unexpected outcome: {:?}", + result.payment_outcome + ); + assert_eq!(result.htlc_count, MAX_RETRIES as usize + 1); + } + fn create_intercept_request(shutdown_listener: Listener) -> InterceptRequest { let (_, pubkey) = get_random_keypair(); InterceptRequest { @@ -2940,14 +3397,21 @@ mod tests { let mock_1 = Arc::new(mock_interceptor_1); let mut test_kit = DispatchPaymentTestKit::new(500_000_000, vec![mock_1], CustomRecords::default()).await; - let (_, result) = test_kit + let (route, result) = test_kit .send_test_payment(test_kit.nodes[0], test_kit.nodes[3], 150_000_000) .await; + // The interceptor acting for Bob rejected the htlc, so it was only added on Alice --> Bob and the channel + // that Bob would have forwarded on is blamed. + let outcome = result.unwrap(); assert!(matches!( - result.unwrap().payment_outcome, + outcome.result.payment_outcome, PaymentOutcome::IndexFailure(0) )); + assert_eq!( + outcome.failed_channel, + Some(route.paths[0].hops[1].short_channel_id) + ); // The interceptor returned a forwarding error, check that a simulation shutdown has not // been triggered.