use std::{
collections::HashMap,
fmt::Debug,
sync::{
atomic::{AtomicU64, Ordering},
Arc, Mutex,
},
};
use linera_base::{data_types::Timestamp, time::Duration};
use tokio::sync::broadcast;
use super::{
cache::SubsumingKey,
request::{RequestKey, RequestResult},
};
use crate::node::NodeError;
#[derive(Debug, Clone)]
pub(super) struct InFlightTracker<N> {
entries: Arc<Mutex<HashMap<RequestKey, InFlightEntry<N>>>>,
next_generation: Arc<AtomicU64>,
timeout: Duration,
}
impl<N: Clone> InFlightTracker<N> {
pub(super) fn new(timeout: Duration) -> Self {
Self {
entries: Arc::new(Mutex::new(HashMap::new())),
next_generation: Arc::new(AtomicU64::new(0)),
timeout,
}
}
pub(super) fn try_subscribe(&self, key: &RequestKey, now: Timestamp) -> Option<InFlightMatch> {
let in_flight = self
.entries
.lock()
.expect("in-flight tracker mutex poisoned");
if let Some(entry) = in_flight.get(key) {
let elapsed = now.duration_since(entry.started_at);
if elapsed <= self.timeout {
return Some(InFlightMatch::Exact(Subscribed(entry.sender.subscribe())));
}
}
for (in_flight_key, entry) in in_flight.iter() {
if in_flight_key.subsumes(key) {
let elapsed = now.duration_since(entry.started_at);
if elapsed <= self.timeout {
return Some(InFlightMatch::Subsuming {
key: in_flight_key.clone(),
outcome: Subscribed(entry.sender.subscribe()),
});
}
}
}
None
}
pub(super) fn insert_new(&self, key: RequestKey, now: Timestamp) -> InFlightGuard<N> {
let generation = self.next_generation.fetch_add(1, Ordering::Relaxed);
let mut in_flight = self
.entries
.lock()
.expect("in-flight tracker mutex poisoned");
let sender = match in_flight.remove(&key) {
Some(previous) => previous.sender,
None => broadcast::channel(1).0,
};
in_flight.insert(
key.clone(),
InFlightEntry {
sender,
started_at: now,
generation,
alternative_peers: Arc::new(tokio::sync::RwLock::new(Vec::new())),
},
);
drop(in_flight);
InFlightGuard {
entries: self.entries.clone(),
key,
generation,
}
}
pub(super) async fn add_alternative_peer(&self, key: &RequestKey, peer: N)
where
N: PartialEq + Eq,
{
let Some(alternative_peers) = self.alternative_peers(key) else {
return;
};
let mut alt_peers = alternative_peers.write().await;
if !alt_peers.contains(&peer) {
alt_peers.push(peer);
}
}
pub(super) async fn get_alternative_peers(&self, key: &RequestKey) -> Option<Vec<N>> {
let alternative_peers = self.alternative_peers(key)?;
let peers = alternative_peers.read().await;
Some(peers.clone())
}
pub(super) async fn remove_alternative_peer(&self, key: &RequestKey, peer: &N)
where
N: PartialEq + Eq,
{
let Some(alternative_peers) = self.alternative_peers(key) else {
return;
};
alternative_peers.write().await.retain(|p| p != peer);
}
pub(super) async fn pop_alternative_peer(&self, key: &RequestKey) -> Option<N> {
self.alternative_peers(key)?.write().await.pop()
}
fn alternative_peers(&self, key: &RequestKey) -> Option<Arc<tokio::sync::RwLock<Vec<N>>>> {
self.entries
.lock()
.expect("in-flight tracker mutex poisoned")
.get(key)
.map(|entry| entry.alternative_peers.clone())
}
}
#[derive(Debug)]
pub(super) enum InFlightMatch {
Exact(Subscribed),
Subsuming {
key: RequestKey,
outcome: Subscribed,
},
}
#[derive(Debug)]
pub(super) struct Subscribed(pub(super) broadcast::Receiver<Arc<Result<RequestResult, NodeError>>>);
pub(super) struct InFlightGuard<N> {
entries: Arc<Mutex<HashMap<RequestKey, InFlightEntry<N>>>>,
key: RequestKey,
generation: u64,
}
impl<N> InFlightGuard<N> {
pub(super) fn complete_and_broadcast(
&self,
result: Arc<Result<RequestResult, NodeError>>,
) -> usize {
let mut in_flight = self
.entries
.lock()
.expect("in-flight tracker mutex poisoned");
if in_flight
.get(&self.key)
.is_none_or(|entry| entry.generation != self.generation)
{
return 0;
}
let entry = in_flight.remove(&self.key).expect("checked above");
let waiter_count = entry.sender.receiver_count();
tracing::trace!(
key = ?self.key,
waiters = waiter_count,
"request completed; broadcasting result to waiters",
);
if waiter_count != 0 {
if let Err(err) = entry.sender.send(result) {
tracing::warn!(
key = ?self.key,
error = ?err,
"failed to broadcast result to waiters"
);
}
}
waiter_count
}
}
impl<N> Drop for InFlightGuard<N> {
fn drop(&mut self) {
let Ok(mut in_flight) = self.entries.lock() else {
return; };
if in_flight
.get(&self.key)
.is_some_and(|entry| entry.generation == self.generation)
{
in_flight.remove(&self.key);
}
}
}
#[derive(Debug)]
pub(super) struct InFlightEntry<N> {
sender: broadcast::Sender<Arc<Result<RequestResult, NodeError>>>,
started_at: Timestamp,
generation: u64,
alternative_peers: Arc<tokio::sync::RwLock<Vec<N>>>,
}