#![cfg(with_server)]
use std::{
collections::{hash_map::Entry, HashMap},
future::Future,
panic::AssertUnwindSafe,
time::Duration,
};
use futures::{channel::mpsc, FutureExt as _, StreamExt as _};
use linera_base::identifiers::ChainId;
#[cfg(with_metrics)]
use linera_base::time::Instant;
use linera_core::data_types::CrossChainRequest;
use rand::Rng as _;
use tracing::{error, trace, warn};
use crate::config::ShardId;
#[cfg(with_metrics)]
pub(crate) mod metrics {
use linera_base::prometheus_util::{
exponential_bucket_latencies, register_histogram, register_int_gauge,
};
use prometheus::{Histogram, IntGauge};
linera_base::declare_metrics! {
pub static CROSS_CHAIN_MESSAGE_TASKS: IntGauge =
register_int_gauge(
"cross_chain_message_tasks",
"Number of concurrent cross-chain message tasks",
);
pub static CROSS_CHAIN_QUEUE_WAIT_TIME: Histogram =
register_histogram(
"cross_chain_queue_wait_time",
"Time (ms) a cross-chain message waits in queue before handle_request is called",
exponential_bucket_latencies(10_000.0),
);
}
}
#[expect(clippy::too_many_arguments)]
pub(crate) async fn forward_cross_chain_queries<F, G>(
nickname: String,
cross_chain_max_retries: u32,
cross_chain_retry_delay: Duration,
cross_chain_max_backoff: Duration,
cross_chain_sender_delay: Duration,
cross_chain_sender_failure_rate: f32,
this_shard: ShardId,
mut receiver: mpsc::Receiver<(CrossChainRequest, ShardId)>,
handle_request: F,
) where
F: Fn(ShardId, CrossChainRequest) -> G + Send + Clone + 'static,
G: Future<Output = anyhow::Result<()>>,
{
let mut steps = futures::stream::FuturesUnordered::new();
let mut job_states: HashMap<QueueId, JobState> = HashMap::new();
let run_task = |task: Task| async move {
#[cfg(with_metrics)]
{
let queue_wait_time_ms = task.queued_at.elapsed().as_secs_f64() * 1000.0;
metrics::CROSS_CHAIN_QUEUE_WAIT_TIME.observe(queue_wait_time_ms);
}
handle_request(task.shard_id, task.request).await
};
let run_action = |action, queue, state: JobState| async move {
linera_base::time::timer::sleep(cross_chain_sender_delay).await;
let to_shard = state.task.shard_id;
(
queue,
match action {
Action::Proceed { .. } => {
let target_chain_id = state.task.request.target_chain_id();
if let Err(error) = run_task(state.task).await {
warn!(
nickname = state.nickname,
?error,
retry = state.retries,
from_shard = this_shard,
to_shard,
chain_id = %target_chain_id,
"Failed to send cross-chain query",
);
Action::Retry
} else {
trace!(from_shard = this_shard, to_shard, "Sent cross-chain query",);
Action::Proceed {
id: state.id.wrapping_add(1),
}
}
}
Action::Retry => {
let delay = cross_chain_retry_delay
.saturating_mul(state.retries)
.min(cross_chain_max_backoff);
linera_base::time::timer::sleep(delay).await;
Action::Proceed { id: state.id }
}
},
)
};
let run_action = move |action, queue: QueueId, state: JobState| {
let nickname = state.nickname.clone();
let to_shard = state.task.shard_id;
let retries = state.retries;
let step = run_action.clone()(action, queue, state);
async move {
AssertUnwindSafe(step)
.catch_unwind()
.await
.unwrap_or_else(|_| {
error!(
nickname,
retry = retries,
from_shard = this_shard,
to_shard,
sender = %queue.sender,
recipient = %queue.recipient,
"Panic while sending a cross-chain query; treating it as a failed attempt",
);
(queue, Action::Retry)
})
}
};
loop {
#[cfg(with_metrics)]
metrics::CROSS_CHAIN_MESSAGE_TASKS.set(job_states.len() as i64);
tokio::select! {
Some((queue, action)) = steps.next() => {
let Entry::Occupied(mut state) = job_states.entry(queue) else {
panic!("running job without state");
};
if state.get().is_finished(&action, cross_chain_max_retries) {
state.remove();
continue;
}
if let Action::Retry = action {
state.get_mut().retries += 1
}
steps.push(run_action.clone()(action, queue, state.get().clone()));
}
request = receiver.next() => {
let Some((request, shard_id)) = request else { break };
if rand::thread_rng().gen::<f32>() < cross_chain_sender_failure_rate {
warn!("Dropped 1 cross-chain message intentionally.");
continue;
}
let queue = QueueId::new(&request);
let task = Task {
shard_id,
request,
#[cfg(with_metrics)]
queued_at: Instant::now(),
};
match job_states.entry(queue) {
Entry::Vacant(entry) => steps.push(run_action.clone()(
Action::Proceed { id: 0 },
queue,
entry.insert(JobState {
id: 0,
retries: 0,
nickname: nickname.clone(),
task,
}).clone(),
)),
Entry::Occupied(mut entry) => {
entry.insert(JobState {
id: entry.get().id + 1,
retries: 0,
nickname: nickname.clone(),
task,
});
}
}
}
else => (),
}
}
}
#[derive(Copy, Clone, PartialEq, Eq, Hash)]
struct QueueId {
sender: ChainId,
recipient: ChainId,
is_update: bool,
}
impl QueueId {
fn new(request: &CrossChainRequest) -> Self {
let (sender, recipient, is_update) = match request {
CrossChainRequest::UpdateRecipient {
sender, recipient, ..
} => (*sender, *recipient, true),
CrossChainRequest::ConfirmUpdatedRecipient {
sender, recipient, ..
}
| CrossChainRequest::RevertConfirm {
sender, recipient, ..
} => (*sender, *recipient, false),
};
QueueId {
sender,
recipient,
is_update,
}
}
}
enum Action {
Proceed { id: usize },
Retry,
}
#[derive(Clone)]
struct Task {
pub shard_id: ShardId,
pub request: linera_core::data_types::CrossChainRequest,
#[cfg(with_metrics)]
pub queued_at: Instant,
}
#[derive(Clone)]
struct JobState {
pub id: usize,
pub retries: u32,
pub nickname: String,
pub task: Task,
}
impl JobState {
fn is_finished(&self, action: &Action, max_retries: u32) -> bool {
match action {
Action::Proceed { id } => self.id < *id,
Action::Retry => self.retries >= max_retries,
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc, Mutex,
};
use futures::{future::BoxFuture, SinkExt as _};
use linera_base::{crypto::CryptoHash, data_types::BlockHeight, identifiers::ChainId};
use tokio::{
sync::{
mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender},
Mutex as AsyncMutex, Semaphore,
},
task::JoinHandle,
};
use super::*;
const NO_DELAY: Duration = Duration::ZERO;
const NEVER_DROP: f32 = 0.0;
const ALWAYS_DROP: f32 = 1.0;
const THIS_SHARD: ShardId = 7;
const SETTLE: Duration = Duration::from_secs(60);
type RequestSender = mpsc::Sender<(CrossChainRequest, ShardId)>;
async fn settle() {
linera_base::time::timer::sleep(SETTLE).await;
}
fn chain(index: u8) -> ChainId {
ChainId(CryptoHash::test_hash(format!("chain {index}")))
}
fn confirm(sender: u8, recipient: u8, tag: u64) -> CrossChainRequest {
CrossChainRequest::ConfirmUpdatedRecipient {
sender: chain(sender),
recipient: chain(recipient),
latest_height: BlockHeight(tag),
}
}
fn update(sender: u8, recipient: u8) -> CrossChainRequest {
CrossChainRequest::UpdateRecipient {
sender: chain(sender),
recipient: chain(recipient),
bundles: Vec::new(),
previous_height: None,
}
}
#[derive(Clone)]
struct Handler {
calls: Arc<Mutex<Vec<(ShardId, CrossChainRequest)>>>,
failures_left: Arc<AtomicUsize>,
panics_left: Arc<AtomicUsize>,
gate: Option<Arc<Semaphore>>,
call_signal: UnboundedSender<()>,
call_signals: Arc<AsyncMutex<UnboundedReceiver<()>>>,
}
impl Handler {
fn new() -> Self {
let (call_signal, call_signals) = unbounded_channel();
Handler {
calls: Arc::default(),
failures_left: Arc::new(AtomicUsize::new(0)),
panics_left: Arc::new(AtomicUsize::new(0)),
gate: None,
call_signal,
call_signals: Arc::new(AsyncMutex::new(call_signals)),
}
}
fn failing(self, count: usize) -> Self {
self.failures_left.store(count, Ordering::SeqCst);
self
}
fn panicking(self, count: usize) -> Self {
self.panics_left.store(count, Ordering::SeqCst);
self
}
fn gated(mut self) -> Self {
self.gate = Some(Arc::new(Semaphore::new(0)));
self
}
fn release(&self, count: usize) {
self.gate
.as_ref()
.expect("handler is gated")
.add_permits(count);
}
fn as_fn(
&self,
) -> impl Fn(ShardId, CrossChainRequest) -> BoxFuture<'static, anyhow::Result<()>>
+ Send
+ Clone
+ 'static {
let this = self.clone();
move |shard_id, request| {
let this = this.clone();
Box::pin(async move { this.call(shard_id, request).await })
}
}
async fn call(self, shard_id: ShardId, request: CrossChainRequest) -> anyhow::Result<()> {
self.calls.lock().unwrap().push((shard_id, request));
self.call_signal.send(()).ok();
if let Some(gate) = &self.gate {
gate.acquire()
.await
.expect("the gate is never closed")
.forget();
}
let take = |budget: &AtomicUsize| {
budget
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |left| {
left.checked_sub(1)
})
.is_ok()
};
if take(&self.panics_left) {
panic!("simulated transport panic");
}
if take(&self.failures_left) {
anyhow::bail!("simulated transport failure");
}
Ok(())
}
async fn wait_for_calls(&self, count: usize) {
let wait = async {
let mut signals = self.call_signals.lock().await;
while self.call_count() < count {
signals.recv().await.expect("the handler outlives the test");
}
};
linera_base::time::timer::timeout(SETTLE, wait)
.await
.unwrap_or_else(|_| {
panic!(
"expected {count} calls, but the forwarder stopped after {}",
self.call_count()
)
});
}
fn calls(&self) -> Vec<(ShardId, CrossChainRequest)> {
self.calls.lock().unwrap().clone()
}
fn call_count(&self) -> usize {
self.calls.lock().unwrap().len()
}
fn tags(&self) -> Vec<u64> {
self.calls()
.into_iter()
.map(|(_, request)| match request {
CrossChainRequest::ConfirmUpdatedRecipient { latest_height, .. } => {
latest_height.0
}
other => panic!("not a confirmation: {other:?}"),
})
.collect()
}
}
fn spawn_forwarder(
handler: &Handler,
max_retries: u32,
retry_delay: Duration,
max_backoff: Duration,
sender_delay: Duration,
failure_rate: f32,
) -> (RequestSender, JoinHandle<()>) {
let (sender, receiver) = mpsc::channel(100);
let task = tokio::spawn(forward_cross_chain_queries(
"test".to_string(),
max_retries,
retry_delay,
max_backoff,
sender_delay,
failure_rate,
THIS_SHARD,
receiver,
handler.as_fn(),
));
(sender, task)
}
fn spawn_simple_forwarder(handler: &Handler) -> (RequestSender, JoinHandle<()>) {
spawn_forwarder(handler, 0, NO_DELAY, NO_DELAY, NO_DELAY, NEVER_DROP)
}
async fn send(sender: &mut RequestSender, request: CrossChainRequest, shard_id: ShardId) {
sender
.send((request, shard_id))
.await
.expect("the forwarder is running");
}
#[tokio::test(start_paused = true)]
async fn test_forwards_a_request_to_its_target_shard() {
let handler = Handler::new();
let (mut sender, _task) = spawn_simple_forwarder(&handler);
send(&mut sender, confirm(1, 2, 5), 3).await;
settle().await;
assert_eq!(handler.calls(), vec![(3, confirm(1, 2, 5))]);
}
#[tokio::test(start_paused = true)]
async fn test_preserves_order_within_a_queue() {
let handler = Handler::new();
let (mut sender, _task) = spawn_simple_forwarder(&handler);
for tag in 1..=3 {
send(&mut sender, confirm(1, 2, tag), 0).await;
settle().await;
}
assert_eq!(handler.tags(), vec![1, 2, 3]);
}
#[tokio::test(start_paused = true)]
async fn test_coalesces_requests_queued_behind_an_in_flight_one() {
let handler = Handler::new().gated();
let (mut sender, _task) = spawn_simple_forwarder(&handler);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1], "the first request is in flight");
send(&mut sender, confirm(1, 2, 2), 0).await;
send(&mut sender, confirm(1, 2, 3), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1], "nothing else has started yet");
handler.release(1);
settle().await;
assert_eq!(handler.tags(), vec![1, 3], "request 2 was superseded by 3");
}
#[tokio::test(start_paused = true)]
async fn test_queues_for_different_recipients_are_independent() {
let handler = Handler::new().gated();
let (mut sender, _task) = spawn_simple_forwarder(&handler);
send(&mut sender, confirm(1, 2, 1), 0).await;
send(&mut sender, confirm(1, 3, 2), 0).await;
settle().await;
let mut tags = handler.tags();
tags.sort_unstable();
assert_eq!(tags, vec![1, 2], "both queues started concurrently");
}
#[tokio::test(start_paused = true)]
async fn test_updates_and_confirmations_are_separate_queues() {
let handler = Handler::new().gated();
let (mut sender, _task) = spawn_simple_forwarder(&handler);
send(&mut sender, update(1, 2), 0).await;
send(&mut sender, confirm(1, 2, 9), 0).await;
settle().await;
assert_eq!(
handler.call_count(),
2,
"the update did not hold up the confirmation",
);
}
#[tokio::test(start_paused = true)]
async fn test_retries_until_the_request_succeeds() {
let handler = Handler::new().failing(2);
let (mut sender, _task) = spawn_forwarder(
&handler,
5,
Duration::from_secs(1),
Duration::from_secs(30),
NO_DELAY,
NEVER_DROP,
);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1, 1, 1], "two failures, then success");
}
#[tokio::test(start_paused = true)]
async fn test_gives_up_after_max_retries() {
let handler = Handler::new().failing(usize::MAX);
let (mut sender, _task) = spawn_forwarder(
&handler,
2,
Duration::from_secs(1),
Duration::from_secs(30),
NO_DELAY,
NEVER_DROP,
);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.call_count(), 3, "one attempt plus two retries");
}
#[tokio::test(start_paused = true)]
async fn test_does_not_retry_when_max_retries_is_zero() {
let handler = Handler::new().failing(usize::MAX);
let (mut sender, _task) = spawn_simple_forwarder(&handler);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.call_count(), 1);
}
#[tokio::test(start_paused = true)]
async fn test_queue_recovers_after_giving_up_on_a_request() {
let handler = Handler::new().failing(3);
let (mut sender, _task) = spawn_forwarder(
&handler,
2,
Duration::from_secs(1),
Duration::from_secs(30),
NO_DELAY,
NEVER_DROP,
);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(
handler.tags(),
vec![1, 1, 1],
"gave up on the first request",
);
send(&mut sender, confirm(1, 2, 2), 0).await;
settle().await;
assert_eq!(
handler.tags(),
vec![1, 1, 1, 2],
"the queue is usable again",
);
}
#[tokio::test(start_paused = true)]
async fn test_retries_after_a_panicking_transport() {
let handler = Handler::new().panicking(1);
let (mut sender, _task) = spawn_forwarder(
&handler,
5,
Duration::from_secs(1),
Duration::from_secs(30),
NO_DELAY,
NEVER_DROP,
);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1, 1], "one panic, then success");
}
#[tokio::test(start_paused = true)]
async fn test_queue_recovers_after_giving_up_on_a_panicking_request() {
let handler = Handler::new().panicking(2);
let (mut sender, _task) = spawn_forwarder(
&handler,
1,
Duration::from_secs(1),
Duration::from_secs(30),
NO_DELAY,
NEVER_DROP,
);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1, 1], "one attempt plus one retry");
send(&mut sender, confirm(1, 2, 2), 0).await;
settle().await;
assert_eq!(handler.tags(), vec![1, 1, 2], "the queue is usable again");
}
#[tokio::test(start_paused = true)]
async fn test_retry_delay_grows_linearly_up_to_the_maximum() {
let handler = Handler::new().failing(usize::MAX);
let (mut sender, _task) = spawn_forwarder(
&handler,
4,
Duration::from_secs(10),
Duration::from_secs(25),
NO_DELAY,
NEVER_DROP,
);
let start = tokio::time::Instant::now();
send(&mut sender, confirm(1, 2, 1), 0).await;
let mut attempt_times = Vec::new();
for attempt in 1..=5 {
handler.wait_for_calls(attempt).await;
attempt_times.push(start.elapsed());
}
settle().await;
assert_eq!(handler.call_count(), 5, "one attempt plus four retries");
assert_eq!(
attempt_times,
vec![
Duration::ZERO,
Duration::from_secs(10),
Duration::from_secs(30),
Duration::from_secs(55),
Duration::from_secs(80),
],
);
}
#[tokio::test(start_paused = true)]
async fn test_sender_delay_postpones_every_request() {
let sender_delay = Duration::from_secs(5);
let handler = Handler::new();
let (mut sender, _task) =
spawn_forwarder(&handler, 0, NO_DELAY, NO_DELAY, sender_delay, NEVER_DROP);
let start = tokio::time::Instant::now();
send(&mut sender, confirm(1, 2, 1), 0).await;
handler.wait_for_calls(1).await;
assert_eq!(start.elapsed(), sender_delay);
}
#[tokio::test(start_paused = true)]
async fn test_failure_rate_of_one_drops_every_request() {
let handler = Handler::new();
let (mut sender, _task) =
spawn_forwarder(&handler, 10, NO_DELAY, NO_DELAY, NO_DELAY, ALWAYS_DROP);
for tag in 1..=5 {
send(&mut sender, confirm(1, 2, tag), 0).await;
}
settle().await;
assert_eq!(handler.call_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn test_closing_the_channel_stops_the_forwarder() {
let handler = Handler::new();
let (mut sender, task) = spawn_simple_forwarder(&handler);
send(&mut sender, confirm(1, 2, 1), 0).await;
settle().await;
drop(sender);
linera_base::time::timer::timeout(SETTLE, task)
.await
.expect("the forwarder stops once the channel is closed")
.expect("the forwarder does not panic");
}
}