use crate::{
marshal::core::{durability::Durable as _, Mailbox, Variant},
types::Round,
};
use commonware_cryptography::{certificate::Scheme, Digest};
use commonware_macros::select;
use commonware_runtime::Handle;
use commonware_utils::{
channel::{fallible::OneshotExt, oneshot},
sync::Mutex,
};
use std::{collections::HashMap, future::Future, sync::Arc};
use tracing::debug;
type Staged<B> = (Arc<B>, oneshot::Sender<Handle<()>>);
struct Inner<D: Digest, B> {
certifications: HashMap<(Round, D), oneshot::Receiver<bool>>,
proposals: HashMap<(Round, D), Staged<B>>,
}
#[derive(Clone)]
pub(crate) struct Gates<D: Digest, B> {
inner: Arc<Mutex<Inner<D, B>>>,
}
impl<D: Digest, B> Default for Gates<D, B> {
fn default() -> Self {
Self::new()
}
}
impl<D: Digest, B> Gates<D, B> {
pub(crate) fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(Inner {
certifications: HashMap::new(),
proposals: HashMap::new(),
})),
}
}
pub(crate) fn insert(&self, round: Round, digest: D, task: oneshot::Receiver<bool>) {
self.inner
.lock()
.certifications
.insert((round, digest), task);
}
pub(crate) fn take(&self, round: Round, digest: D) -> Option<oneshot::Receiver<bool>> {
self.inner.lock().certifications.remove(&(round, digest))
}
pub(crate) fn take_staged(&self, round: Round, digest: D) -> Option<Staged<B>> {
self.inner.lock().proposals.remove(&(round, digest))
}
pub(crate) fn flush_unrelayed<S, V>(&self, marshal: &Mailbox<S, V>, round: Round, id: D)
where
S: Scheme,
V: Variant<Block = B>,
{
if let Some((block, ack)) = self.take_staged(round, id) {
marshal.verified_deferred(round, block, ack);
}
}
pub(crate) fn retain_after(&self, finalized_round: &Round) {
let mut inner = self.inner.lock();
inner
.certifications
.retain(|(round, _), _| round > finalized_round);
inner
.proposals
.retain(|(round, _), _| round > finalized_round);
}
pub(crate) async fn stage(
&self,
round: Round,
id: D,
block: Arc<B>,
tx: oneshot::Sender<D>,
name: &'static str,
) {
let (durable_tx, durable_rx) = oneshot::channel();
let (ack, persist) = oneshot::channel();
{
let mut inner = self.inner.lock();
inner.certifications.insert((round, id), durable_rx);
inner.proposals.insert((round, id), (block, ack));
}
tx.send_lossy(id);
let Ok(handle) = persist.await else {
return;
};
if !handle.durable(round, name).await {
return;
}
durable_tx.send_lossy(true);
debug!(?round, ?id, name, "block durable");
}
}
pub(crate) const fn resolve(verdict: Option<bool>, durable: bool) -> Option<bool> {
match verdict {
Some(true) if !durable => None,
other => other,
}
}
pub(crate) async fn drive<D, F, Fut>(
mut tx: oneshot::Sender<bool>,
task: oneshot::Receiver<bool>,
round: Round,
id: D,
fallback: F,
) where
D: Digest,
F: FnOnce() -> Fut,
Fut: Future<Output = oneshot::Receiver<bool>>,
{
let result = select! {
_ = tx.closed() => {
debug!(reason = "consensus dropped receiver", "skipping certification");
return;
},
result = task => result,
};
match result {
Ok(result) => {
tx.send_lossy(result);
}
Err(_) => {
debug!(
?round,
?id,
"certification gate task closed before certification, falling back to embedded context"
);
let fallback = fallback().await;
let result = select! {
_ = tx.closed() => {
debug!(reason = "consensus dropped receiver", "skipping certification");
return;
},
result = fallback => result,
};
if let Ok(result) = result {
tx.send_lossy(result);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{Epoch, View};
use commonware_cryptography::{sha256::Digest as Sha256Digest, Hasher, Sha256};
use commonware_runtime::{deterministic, Runner, Spawner};
type D = Sha256Digest;
type TestGates = Gates<D, u64>;
fn round(view: u64) -> Round {
Round::new(Epoch::zero(), View::new(view))
}
fn pending_task() -> oneshot::Receiver<bool> {
let (_tx, rx) = oneshot::channel();
rx
}
#[test]
fn test_insert_and_take_returns_task() {
let tasks = TestGates::new();
let digest = Sha256::hash(b"block");
tasks.insert(round(1), digest, pending_task());
assert!(tasks.take(round(1), digest).is_some());
assert!(
tasks.take(round(1), digest).is_none(),
"taking twice should yield None"
);
}
#[test]
fn test_take_absent_key_is_none() {
let tasks = TestGates::new();
assert!(tasks.take(round(1), Sha256::hash(b"missing")).is_none());
}
#[test]
fn test_take_distinguishes_rounds_and_digests() {
let tasks = TestGates::new();
let digest_a = Sha256::hash(b"a");
let digest_b = Sha256::hash(b"b");
tasks.insert(round(1), digest_a, pending_task());
tasks.insert(round(2), digest_a, pending_task());
tasks.insert(round(1), digest_b, pending_task());
assert!(tasks.take(round(1), digest_a).is_some());
assert!(tasks.take(round(2), digest_a).is_some());
assert!(tasks.take(round(1), digest_b).is_some());
}
#[test]
fn test_retain_after_drops_at_and_below_boundary() {
let tasks = TestGates::new();
let digest = Sha256::hash(b"block");
tasks.insert(round(1), digest, pending_task());
tasks.insert(round(2), digest, pending_task());
tasks.insert(round(3), digest, pending_task());
tasks.retain_after(&round(2));
assert!(
tasks.take(round(1), digest).is_none(),
"tasks strictly below boundary should be dropped"
);
assert!(
tasks.take(round(2), digest).is_none(),
"tasks at boundary should be dropped"
);
assert!(
tasks.take(round(3), digest).is_some(),
"tasks strictly above boundary should be retained"
);
}
#[test]
fn test_retain_after_spans_epochs() {
let tasks = TestGates::new();
let digest = Sha256::hash(b"block");
let early = Round::new(Epoch::zero(), View::new(100));
let late = Round::new(Epoch::new(1), View::zero());
tasks.insert(early, digest, pending_task());
tasks.insert(late, digest, pending_task());
tasks.retain_after(&early);
assert!(
tasks.take(early, digest).is_none(),
"task at boundary must be dropped"
);
assert!(
tasks.take(late, digest).is_some(),
"task in later epoch must outlive an earlier boundary"
);
}
#[test]
fn test_retain_after_empty_map_is_noop() {
let tasks = TestGates::new();
tasks.retain_after(&round(5));
assert!(tasks.take(round(5), Sha256::hash(b"x")).is_none());
}
#[test]
fn test_default_matches_new() {
let default = <TestGates as Default>::default();
let digest = Sha256::hash(b"block");
default.insert(round(1), digest, pending_task());
assert!(default.take(round(1), digest).is_some());
}
#[test]
fn test_resolve() {
assert_eq!(resolve(None, true), None);
assert_eq!(resolve(None, false), None);
assert_eq!(resolve(Some(false), false), Some(false));
assert_eq!(resolve(Some(false), true), Some(false));
assert_eq!(resolve(Some(true), true), Some(true));
assert_eq!(resolve(Some(true), false), None);
}
#[test]
fn test_stage_handshake() {
let runner = deterministic::Runner::default();
runner.start(|context| async move {
let gates = TestGates::new();
let digest = Sha256::hash(b"block");
let (tx, rx) = oneshot::channel();
context.spawn({
let gates = gates.clone();
move |_| async move {
gates.stage(round(1), digest, Arc::new(7), tx, "test").await;
}
});
assert_eq!(rx.await.expect("id published"), digest);
let gate = gates.take(round(1), digest).expect("gate registered");
let (block, ack) = gates.take_staged(round(1), digest).expect("block staged");
assert_eq!(*block, 7);
assert!(
gates.take_staged(round(1), digest).is_none(),
"taking twice should yield None"
);
ack.send_lossy(Handle::ready(Ok(())));
assert!(gate.await.expect("gate resolved"));
});
}
#[test]
fn test_retain_after_drops_staged_and_abandons_handshake() {
let runner = deterministic::Runner::default();
runner.start(|context| async move {
let gates = TestGates::new();
let digest = Sha256::hash(b"block");
let (tx, rx) = oneshot::channel();
context.spawn({
let gates = gates.clone();
move |_| async move {
gates.stage(round(1), digest, Arc::new(7), tx, "test").await;
}
});
assert_eq!(rx.await.expect("id published"), digest);
let gate = gates.take(round(1), digest).expect("gate registered");
gates.retain_after(&round(1));
assert!(gates.take_staged(round(1), digest).is_none());
assert!(gate.await.is_err(), "gate must be abandoned, not resolved");
});
}
}