use crate::stateful::db::{AttachableResolver, Shared, p2p::cancel};
use commonware_actor::mailbox::{Overflow, Policy, Sender};
use commonware_codec::Read;
use commonware_cryptography::Digest;
use commonware_storage::{
merkle::Family,
qmdb::sync::{FeedbackTx, Request, Response, Source},
};
use commonware_utils::channel::oneshot;
use std::{collections::VecDeque, future::Future};
#[derive(Debug, thiserror::Error)]
#[error("response dropped before completion")]
pub struct ResponseDropped;
pub(super) type ResponseTx<F, Op, D> = oneshot::Sender<(Response<F, Op, D>, FeedbackTx)>;
pub(super) enum Message<DB, F: Family, Op, D: Digest> {
AttachDatabase(Shared<DB>),
GetOperations {
request: Request<F>,
response: ResponseTx<F, Op, D>,
},
CancelOperations { request: Request<F> },
}
impl<DB, F: Family, Op, D: Digest> Message<DB, F, Op, D> {
fn response_closed(&self) -> bool {
match self {
Self::AttachDatabase(_) | Self::CancelOperations { .. } => false,
Self::GetOperations { response, .. } => response.is_closed(),
}
}
}
pub(super) struct Pending<DB, F: Family, Op, D: Digest> {
database: Option<Shared<DB>>,
messages: VecDeque<Message<DB, F, Op, D>>,
}
impl<DB, F: Family, Op, D: Digest> Default for Pending<DB, F, Op, D> {
fn default() -> Self {
Self {
database: None,
messages: VecDeque::new(),
}
}
}
impl<DB, F: Family, Op, D: Digest> Overflow<Message<DB, F, Op, D>> for Pending<DB, F, Op, D> {
fn is_empty(&self) -> bool {
self.database.is_none() && self.messages.is_empty()
}
fn drain<P>(&mut self, mut push: P)
where
P: FnMut(Message<DB, F, Op, D>) -> Option<Message<DB, F, Op, D>>,
{
if let Some(database) = self.database.take()
&& let Some(Message::AttachDatabase(database)) = push(Message::AttachDatabase(database))
{
self.database = Some(database);
return;
}
while let Some(message) = self.messages.pop_front() {
if message.response_closed() {
continue;
}
if let Some(message) = push(message) {
self.messages.push_front(message);
break;
}
}
}
}
impl<DB, F: Family, Op, D: Digest> Policy for Message<DB, F, Op, D> {
type Overflow = Pending<DB, F, Op, D>;
fn handle(overflow: &mut Self::Overflow, message: Self) {
if message.response_closed() {
return;
}
match message {
Self::AttachDatabase(database) => {
overflow.database = Some(database);
}
message => overflow.messages.push_back(message),
}
}
}
pub struct Mailbox<DB, F: Family, Op, D: Digest> {
sender: Sender<Message<DB, F, Op, D>>,
}
impl<DB, F: Family, Op, D: Digest> Clone for Mailbox<DB, F, Op, D> {
fn clone(&self) -> Self {
Self {
sender: self.sender.clone(),
}
}
}
impl<DB, F: Family, Op, D: Digest> Mailbox<DB, F, Op, D> {
pub(super) const fn new(sender: Sender<Message<DB, F, Op, D>>) -> Self {
Self { sender }
}
}
impl<DB: Send + Sync, F: Family, Op: Send, D: Digest> Mailbox<DB, F, Op, D> {
pub fn attach_database(&self, db: Shared<DB>) {
let _ = self.sender.enqueue(Message::AttachDatabase(db));
}
}
impl<DB, F, Op, D> Source for Mailbox<DB, F, Op, D>
where
F: Family,
Op: Read<Cfg = ()> + Send + Sync + Clone + 'static,
D: Digest,
DB: Send + Sync + 'static,
{
type Family = F;
type Digest = D;
type Op = Op;
type Error = ResponseDropped;
async fn serve(
&self,
request: Request<F>,
) -> Result<(Response<Self::Family, Self::Op, Self::Digest>, FeedbackTx), Self::Error> {
let (response_tx, response_rx) = oneshot::channel();
let _ = self.sender.enqueue(Message::GetOperations {
request,
response: response_tx,
});
let mut guard =
cancel::Guard::new(self.sender.clone(), Message::CancelOperations { request });
let result = response_rx.await;
guard.disarm();
result.map_err(|_| ResponseDropped)
}
}
impl<DB, F, Op, D> AttachableResolver<DB> for Mailbox<DB, F, Op, D>
where
F: Family,
Op: Read<Cfg = ()> + Send + Sync + Clone + 'static,
D: Digest,
DB: Send + Sync + 'static,
{
fn attach_database(&self, db: Shared<DB>) -> impl Future<Output = ()> + Send {
Self::attach_database(self, db);
std::future::ready(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use commonware_cryptography::sha256;
use commonware_runtime::{Runner as _, deterministic};
use commonware_storage::mmr;
use commonware_utils::{NZU64, NZUsize};
#[test]
fn dropping_get_operations_sends_cancel_message() {
deterministic::Runner::default().start(|context| async move {
let (sender, mut receiver) = commonware_actor::mailbox::new(context, NZUsize!(4));
let mailbox = Mailbox::<(), mmr::Family, u64, sha256::Digest>::new(sender);
let size = mmr::Location::new(10);
let start_loc = mmr::Location::new(3);
let max_ops = NZU64!(2);
{
let get = mailbox.serve(Request::Operations {
size,
start: start_loc,
max_ops,
});
futures::pin_mut!(get);
assert!(futures::poll!(get.as_mut()).is_pending());
}
match receiver.recv().await.expect("request should be queued") {
Message::GetOperations { request, .. } => {
assert_eq!(request.size(), size);
assert_eq!(request.start(), start_loc);
assert_eq!(request.max_ops(), max_ops);
assert!(matches!(request, Request::Operations { .. }));
}
Message::AttachDatabase(_) => panic!("unexpected attach message"),
Message::CancelOperations { .. } => panic!("cancel should come after request"),
}
match receiver.recv().await.expect("cancel should be queued") {
Message::CancelOperations { request } => {
assert_eq!(request.size(), size);
assert_eq!(request.start(), start_loc);
assert_eq!(request.max_ops(), max_ops);
assert!(matches!(request, Request::Operations { .. }));
}
Message::AttachDatabase(_) => panic!("unexpected attach message"),
Message::GetOperations { .. } => panic!("unexpected duplicate request"),
}
});
}
#[test]
fn completed_get_operations_sends_no_cancel() {
deterministic::Runner::default().start(|context| async move {
let (sender, mut receiver) = commonware_actor::mailbox::new(context, NZUsize!(4));
let mailbox = Mailbox::<(), mmr::Family, u64, sha256::Digest>::new(sender);
let get = mailbox.serve(Request::Operations {
size: mmr::Location::new(10),
start: mmr::Location::new(3),
max_ops: NZU64!(2),
});
let observe = async move {
let Message::GetOperations { response, .. } =
receiver.recv().await.expect("request should be queued")
else {
panic!("expected a fetch request");
};
drop(response);
receiver
};
let (result, mut receiver) = futures::join!(get, observe);
assert!(matches!(result, Err(ResponseDropped)));
assert!(
receiver.try_recv().is_err(),
"a completed fetch must not enqueue a cancel"
);
});
}
}