use super::executor_spawn::try_spawn;
use super::finalization::EventSequencer;
use monoloop_contracts::{
ChannelId, EventDeliveryError, SessionId, TransactionEvent, TransactionEventPayload,
TransactionEventSink, TransactionId,
};
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::runtime::Handle;
use tokio::sync::{mpsc, Mutex};
pub struct QueuedEvent {
pub event: TransactionEvent,
pub ack: Option<tokio::sync::oneshot::Sender<Result<(), EventDeliveryError>>>,
pub approx_bytes: usize,
}
impl QueuedEvent {
pub fn new(
event: TransactionEvent,
ack: Option<tokio::sync::oneshot::Sender<Result<(), EventDeliveryError>>>,
) -> Self {
let approx_bytes = estimate_event_bytes(&event);
Self {
event,
ack,
approx_bytes,
}
}
}
fn estimate_event_bytes(event: &TransactionEvent) -> usize {
serde_json::to_vec(event)
.map(|b| b.len().max(64))
.unwrap_or(256)
}
#[derive(Clone)]
pub struct BoundedEventSender {
tx: mpsc::Sender<QueuedEvent>,
queued_bytes: Arc<AtomicUsize>,
max_bytes: usize,
}
impl BoundedEventSender {
pub fn new(tx: mpsc::Sender<QueuedEvent>, max_bytes: usize) -> Self {
Self {
tx,
queued_bytes: Arc::new(AtomicUsize::new(0)),
max_bytes: max_bytes.max(1),
}
}
pub async fn send(&self, item: QueuedEvent) -> Result<(), EventQueueFull> {
let bytes = item.approx_bytes;
loop {
let cur = self.queued_bytes.load(Ordering::SeqCst);
if cur.saturating_add(bytes) > self.max_bytes {
return Err(EventQueueFull::Bytes);
}
if self
.queued_bytes
.compare_exchange(cur, cur + bytes, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
break;
}
}
let mut reservation = ByteReservation {
counter: &self.queued_bytes,
bytes,
released: false,
};
match self.tx.send(item).await {
Ok(()) => {
reservation.released = true;
Ok(())
}
Err(_) => {
reservation.release();
Err(EventQueueFull::Closed)
}
}
}
pub fn byte_counter(&self) -> Arc<AtomicUsize> {
Arc::clone(&self.queued_bytes)
}
}
#[derive(Clone)]
pub struct OrderedEventPublisher {
order: Arc<Mutex<()>>,
event_tx: BoundedEventSender,
sequencer: Arc<EventSequencer>,
}
impl OrderedEventPublisher {
pub fn new(event_tx: BoundedEventSender, sequencer: Arc<EventSequencer>) -> Self {
Self {
order: Arc::new(Mutex::new(())),
event_tx,
sequencer,
}
}
pub fn sequencer(&self) -> &Arc<EventSequencer> {
&self.sequencer
}
pub async fn publish(
&self,
transaction_id: TransactionId,
channel_id: ChannelId,
session_id: SessionId,
payload: TransactionEventPayload,
) -> Result<u64, EventQueueFull> {
self.publish_inner(transaction_id, channel_id, session_id, payload, None)
.await
}
pub async fn publish_terminal(
&self,
transaction_id: TransactionId,
channel_id: ChannelId,
session_id: SessionId,
payload: TransactionEventPayload,
ack: tokio::sync::oneshot::Sender<Result<(), EventDeliveryError>>,
) -> Result<u64, EventQueueFull> {
self.publish_inner(transaction_id, channel_id, session_id, payload, Some(ack))
.await
}
async fn publish_inner(
&self,
transaction_id: TransactionId,
channel_id: ChannelId,
session_id: SessionId,
payload: TransactionEventPayload,
ack: Option<tokio::sync::oneshot::Sender<Result<(), EventDeliveryError>>>,
) -> Result<u64, EventQueueFull> {
let _guard = self.order.lock().await;
let seq = self.sequencer.peek_next();
let event = TransactionEvent {
transaction_id,
channel_id,
session_id,
sequence: seq,
payload,
};
self.event_tx.send(QueuedEvent::new(event, ack)).await?;
let got = self.sequencer.allocate();
debug_assert_eq!(got, seq);
Ok(seq)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EventQueueFull {
Bytes,
Closed,
}
struct ByteReservation<'a> {
counter: &'a AtomicUsize,
bytes: usize,
released: bool,
}
impl ByteReservation<'_> {
fn release(&mut self) {
if !self.released {
self.counter.fetch_sub(self.bytes, Ordering::SeqCst);
self.released = true;
}
}
}
impl Drop for ByteReservation<'_> {
fn drop(&mut self) {
self.release();
}
}
pub fn spawn_delivery_task(
executor: &Handle,
mut rx: mpsc::Receiver<QueuedEvent>,
sink: Arc<dyn TransactionEventSink>,
on_fail: mpsc::Sender<()>,
byte_counter: Arc<AtomicUsize>,
deliver_deadline: Duration,
) -> Result<tokio::task::JoinHandle<()>, ()> {
let executor_child = executor.clone();
try_spawn(executor, async move {
while let Some(item) = rx.recv().await {
let bytes = item.approx_bytes;
let result =
deliver_isolated(&executor_child, &sink, item.event, deliver_deadline).await;
byte_counter.fetch_sub(
bytes.min(byte_counter.load(Ordering::SeqCst)),
Ordering::SeqCst,
);
let ok = result.is_ok();
if let Some(ack) = item.ack {
let _ = ack.send(if ok {
Ok(())
} else {
Err(EventDeliveryError::Failed)
});
}
if !ok {
let _ = on_fail.try_send(());
while let Some(rest) = rx.recv().await {
byte_counter.fetch_sub(
rest.approx_bytes.min(byte_counter.load(Ordering::SeqCst)),
Ordering::SeqCst,
);
if let Some(ack) = rest.ack {
let _ = ack.send(Err(EventDeliveryError::Failed));
}
}
break;
}
}
})
}
async fn deliver_isolated(
executor: &Handle,
sink: &Arc<dyn TransactionEventSink>,
event: TransactionEvent,
deadline: Duration,
) -> Result<(), EventDeliveryError> {
let deliver_fut = catch_unwind(AssertUnwindSafe(|| sink.deliver(event)));
let fut = match deliver_fut {
Ok(f) => f,
Err(_) => return Err(EventDeliveryError::Failed),
};
let handle = match try_spawn(executor, fut) {
Ok(h) => h,
Err(()) => return Err(EventDeliveryError::Failed),
};
let abort = handle.abort_handle();
match tokio::time::timeout(deadline, handle).await {
Ok(Ok(r)) => r,
Ok(Err(_)) => Err(EventDeliveryError::Failed),
Err(_) => {
abort.abort();
Err(EventDeliveryError::Failed)
}
}
}