1use crate::transaction::{TransactionCompletion, TransactionEvent};
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::sync::Arc;
10use thiserror::Error;
11use tokio::sync::{mpsc, oneshot, Notify};
12
13#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub struct DeliveryLimits {
16 pub max_event_items: usize,
18 pub max_event_bytes: usize,
20}
21
22impl DeliveryLimits {
23 pub fn try_new(
25 max_event_items: usize,
26 max_event_bytes: usize,
27 ) -> Result<Self, DeliveryConfigError> {
28 if max_event_items == 0 {
29 return Err(DeliveryConfigError::ZeroItemCapacity);
30 }
31 if max_event_bytes == 0 {
32 return Err(DeliveryConfigError::ZeroByteCapacity);
33 }
34 Ok(Self {
35 max_event_items,
36 max_event_bytes,
37 })
38 }
39}
40
41#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)]
43pub enum DeliveryConfigError {
44 #[error("delivery item capacity must be nonzero")]
46 ZeroItemCapacity,
47 #[error("delivery byte capacity must be nonzero")]
49 ZeroByteCapacity,
50}
51
52#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
54pub enum CompletionPublishResult {
55 Published,
57 ReceiverDropped,
59 InvariantFailed,
61}
62
63struct QueuedBytePermit {
67 queued_bytes: Arc<AtomicU64>,
68 byte_freed: Arc<Notify>,
69 nbytes: u64,
70}
71
72impl Drop for QueuedBytePermit {
73 fn drop(&mut self) {
74 self.queued_bytes.fetch_sub(self.nbytes, Ordering::SeqCst);
75 self.byte_freed.notify_waiters();
76 }
77}
78
79struct QueuedDeliveryEvent {
81 event: TransactionEvent,
82 permit: Option<QueuedBytePermit>,
84}
85
86#[derive(Debug, Clone)]
88pub struct TransactionEventSender {
89 tx: mpsc::Sender<QueuedDeliveryEvent>,
90 queued_bytes: Arc<AtomicU64>,
92 byte_freed: Arc<Notify>,
94 max_event_items: usize,
96 max_event_bytes: usize,
97}
98
99#[derive(Debug)]
101pub struct TransactionEventReceiver {
102 rx: mpsc::Receiver<QueuedDeliveryEvent>,
103}
104
105#[derive(Debug)]
107pub struct TransactionCompletionSender {
108 tx: Option<oneshot::Sender<TransactionCompletion>>,
109}
110
111#[derive(Debug)]
113pub struct TransactionCompletionReceiver {
114 rx: oneshot::Receiver<TransactionCompletion>,
115}
116
117#[derive(Debug)]
119pub struct TransactionDelivery {
120 pub event_tx: TransactionEventSender,
122 pub completion_tx: TransactionCompletionSender,
124}
125
126#[derive(Debug)]
128pub struct TransactionReceiver {
129 pub events: TransactionEventReceiver,
131 pub completion: TransactionCompletionReceiver,
133}
134
135pub fn transaction_delivery(
137 limits: DeliveryLimits,
138) -> Result<(TransactionDelivery, TransactionReceiver), DeliveryConfigError> {
139 let limits = DeliveryLimits::try_new(limits.max_event_items, limits.max_event_bytes)?;
140 let (event_tx, event_rx) = mpsc::channel(limits.max_event_items);
141 let (completion_tx, completion_rx) = oneshot::channel();
142 let queued_bytes = Arc::new(AtomicU64::new(0));
143 let byte_freed = Arc::new(Notify::new());
144 Ok((
145 TransactionDelivery {
146 event_tx: TransactionEventSender {
147 tx: event_tx,
148 queued_bytes,
149 byte_freed,
150 max_event_items: limits.max_event_items,
151 max_event_bytes: limits.max_event_bytes,
152 },
153 completion_tx: TransactionCompletionSender {
154 tx: Some(completion_tx),
155 },
156 },
157 TransactionReceiver {
158 events: TransactionEventReceiver { rx: event_rx },
159 completion: TransactionCompletionReceiver { rx: completion_rx },
160 },
161 ))
162}
163
164impl TransactionEventSender {
165 pub fn max_event_items(&self) -> usize {
167 self.max_event_items
168 }
169
170 pub fn max_event_bytes(&self) -> usize {
172 self.max_event_bytes
173 }
174
175 pub fn queued_bytes(&self) -> u64 {
177 self.queued_bytes.load(Ordering::SeqCst)
178 }
179
180 pub fn try_send(&self, event: TransactionEvent) -> Result<(), EventEnqueueError> {
182 let nbytes = estimate_event_bytes(&event);
183 if nbytes > self.max_event_bytes {
184 return Err(EventEnqueueError::EventTooLarge);
185 }
186 loop {
187 let cur = self.queued_bytes.load(Ordering::SeqCst);
188 let next = cur.saturating_add(nbytes as u64);
189 if next > self.max_event_bytes as u64 {
190 return Err(EventEnqueueError::ByteCapacityExceeded);
191 }
192 if self
193 .queued_bytes
194 .compare_exchange(cur, next, Ordering::SeqCst, Ordering::SeqCst)
195 .is_ok()
196 {
197 let item = QueuedDeliveryEvent {
198 event,
199 permit: Some(QueuedBytePermit {
200 queued_bytes: Arc::clone(&self.queued_bytes),
201 byte_freed: Arc::clone(&self.byte_freed),
202 nbytes: nbytes as u64,
203 }),
204 };
205 match self.tx.try_send(item) {
206 Ok(()) => return Ok(()),
207 Err(mpsc::error::TrySendError::Full(_)) => {
209 return Err(EventEnqueueError::ItemCapacityExceeded);
210 }
211 Err(mpsc::error::TrySendError::Closed(_)) => {
212 return Err(EventEnqueueError::Closed);
213 }
214 }
215 }
216 }
217 }
218
219 pub async fn send(&self, event: TransactionEvent) -> Result<(), EventEnqueueError> {
225 let nbytes = estimate_event_bytes(&event);
226 if nbytes > self.max_event_bytes {
227 return Err(EventEnqueueError::EventTooLarge);
228 }
229 loop {
231 let notified = self.byte_freed.notified();
232 tokio::pin!(notified);
233 let cur = self.queued_bytes.load(Ordering::SeqCst);
234 let next = cur.saturating_add(nbytes as u64);
235 if next > self.max_event_bytes as u64 {
236 notified.await;
237 continue;
238 }
239 if self
240 .queued_bytes
241 .compare_exchange(cur, next, Ordering::SeqCst, Ordering::SeqCst)
242 .is_ok()
243 {
244 break;
245 }
246 }
247 let item = QueuedDeliveryEvent {
248 event,
249 permit: Some(QueuedBytePermit {
250 queued_bytes: Arc::clone(&self.queued_bytes),
251 byte_freed: Arc::clone(&self.byte_freed),
252 nbytes: nbytes as u64,
253 }),
254 };
255 match self.tx.send(item).await {
256 Ok(()) => Ok(()),
257 Err(mpsc::error::SendError(_)) => Err(EventEnqueueError::Closed),
259 }
260 }
261}
262
263impl TransactionEventReceiver {
264 pub async fn recv(&mut self) -> Option<TransactionEvent> {
268 self.rx.recv().await.map(|item| {
269 let QueuedDeliveryEvent { event, permit } = item;
270 drop(permit);
271 event
272 })
273 }
274
275 pub fn try_recv(&mut self) -> Result<TransactionEvent, mpsc::error::TryRecvError> {
277 self.rx.try_recv().map(|item| {
278 let QueuedDeliveryEvent { event, permit } = item;
279 drop(permit);
280 event
281 })
282 }
283}
284
285impl TransactionCompletionSender {
286 pub fn send(mut self, completion: TransactionCompletion) -> CompletionPublishResult {
288 match self.tx.take() {
289 Some(tx) => {
290 if tx.send(completion).is_ok() {
291 CompletionPublishResult::Published
292 } else {
293 CompletionPublishResult::ReceiverDropped
294 }
295 }
296 None => CompletionPublishResult::InvariantFailed,
297 }
298 }
299}
300
301impl TransactionCompletionReceiver {
302 pub async fn recv(self) -> Result<TransactionCompletion, oneshot::error::RecvError> {
304 self.rx.await
305 }
306}
307
308#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)]
310pub enum EventEnqueueError {
311 #[error("event item capacity exceeded")]
313 ItemCapacityExceeded,
314 #[error("event byte capacity exceeded")]
316 ByteCapacityExceeded,
317 #[error("event exceeds delivery byte capacity")]
319 EventTooLarge,
320 #[error("event delivery channel closed")]
322 Closed,
323}
324
325pub fn estimate_event_bytes(event: &TransactionEvent) -> usize {
327 match serde_json::to_vec(event) {
329 Ok(buf) => buf.len().max(64),
330 Err(_) => 256,
331 }
332}
333
334#[cfg(test)]
335mod tests {
336 use super::*;
337 use crate::id::{ChannelId, SessionId, TransactionId};
338 use crate::transaction::{
339 CleanupStatus, TerminalEventDelivery, TransactionDiagnostic, TransactionEndEvent,
340 TransactionEndKind, TransactionEventPayload, TransactionUsage,
341 };
342 use std::sync::atomic::{AtomicU64, Ordering};
343 use std::sync::Arc;
344
345 fn sample_completion() -> TransactionCompletion {
346 TransactionCompletion {
347 end: TransactionEndEvent {
348 transaction_id: TransactionId::generate(),
349 session_id: None,
350 channel_id: ChannelId::try_new("ch").unwrap(),
351 kind: TransactionEndKind::Completed,
352 emitted_events: 1,
353 usage: TransactionUsage::default(),
354 diagnostics: Vec::<TransactionDiagnostic>::new(),
355 },
356 terminal_event_delivery: TerminalEventDelivery::Published,
357 cleanup: CleanupStatus::Complete,
358 }
359 }
360
361 #[test]
362 fn rejects_zero_capacities() {
363 assert!(DeliveryLimits::try_new(0, 1024).is_err());
364 assert!(DeliveryLimits::try_new(8, 0).is_err());
365 }
366
367 #[tokio::test]
368 async fn completion_publish_and_receive() {
369 let (delivery, mut receiver) =
370 transaction_delivery(DeliveryLimits::try_new(8, 64 * 1024).unwrap()).unwrap();
371 let expected = sample_completion();
372 let result = delivery.completion_tx.send(expected.clone());
373 assert_eq!(result, CompletionPublishResult::Published);
374 let got = receiver.completion.recv().await.unwrap();
375 assert_eq!(got, expected);
376 let _ = &mut receiver.events;
377 }
378
379 #[tokio::test]
380 async fn completion_receiver_dropped_is_observable() {
381 let (delivery, receiver) =
382 transaction_delivery(DeliveryLimits::try_new(4, 4096).unwrap()).unwrap();
383 drop(receiver);
384 let result = delivery.completion_tx.send(sample_completion());
385 assert_eq!(result, CompletionPublishResult::ReceiverDropped);
386 }
387
388 fn sample_event(sequence: u64) -> TransactionEvent {
389 TransactionEvent {
390 transaction_id: TransactionId::generate(),
391 channel_id: ChannelId::try_new("ch").unwrap(),
392 session_id: SessionId::try_new("s").unwrap(),
393 sequence,
394 payload: TransactionEventPayload::Diagnostic(TransactionDiagnostic {
395 diagnostic: crate::safe::SafeDiagnostic::try_new_default("internal", Some("x"))
396 .unwrap(),
397 }),
398 }
399 }
400
401 #[tokio::test]
402 async fn event_try_send_item_capacity() {
403 let (delivery, _receiver) =
404 transaction_delivery(DeliveryLimits::try_new(1, 1024 * 1024).unwrap()).unwrap();
405 let ev = sample_event(1);
406 delivery.event_tx.try_send(ev.clone()).unwrap();
407 let err = delivery.event_tx.try_send(ev).unwrap_err();
408 assert_eq!(err, EventEnqueueError::ItemCapacityExceeded);
409 }
410
411 #[tokio::test]
414 async fn event_byte_capacity_recovered_after_receive() {
415 let ev = sample_event(1);
416 let one = estimate_event_bytes(&ev);
417 let (delivery, mut receiver) =
418 transaction_delivery(DeliveryLimits::try_new(4, one).unwrap()).unwrap();
419
420 for i in 0..8u64 {
421 let next = sample_event(i + 1);
422 delivery
423 .event_tx
424 .try_send(next)
425 .unwrap_or_else(|e| panic!("cycle {i} must succeed after receive release: {e:?}"));
426 let _ = receiver
427 .events
428 .recv()
429 .await
430 .expect("receive releases bytes");
431 assert_eq!(
432 delivery.event_tx.queued_bytes(),
433 0,
434 "queued bytes must be zero after receive"
435 );
436 }
437 }
438
439 #[tokio::test]
441 async fn event_byte_capacity_exact_and_plus_one() {
442 let ev = sample_event(1);
443 let one = estimate_event_bytes(&ev);
444 let (delivery, mut receiver) =
445 transaction_delivery(DeliveryLimits::try_new(8, one * 2).unwrap()).unwrap();
446 let a = sample_event(1);
447 let b = sample_event(2);
448 let c = sample_event(3);
449 assert_eq!(estimate_event_bytes(&a), one);
450 assert_eq!(estimate_event_bytes(&b), one);
451 assert_eq!(estimate_event_bytes(&c), one);
452
453 delivery.event_tx.try_send(a).unwrap();
454 delivery.event_tx.try_send(b).unwrap();
455 assert_eq!(delivery.event_tx.queued_bytes(), (one * 2) as u64);
456 let err = delivery.event_tx.try_send(c).unwrap_err();
457 assert_eq!(err, EventEnqueueError::ByteCapacityExceeded);
458
459 let _ = receiver.events.recv().await.unwrap();
460 assert_eq!(delivery.event_tx.queued_bytes(), one as u64);
461 delivery.event_tx.try_send(sample_event(4)).unwrap();
462 assert_eq!(delivery.event_tx.queued_bytes(), (one * 2) as u64);
463 }
464
465 #[tokio::test]
467 async fn event_byte_capacity_released_on_receiver_drop() {
468 let ev = sample_event(1);
469 let one = estimate_event_bytes(&ev);
470 let (delivery, receiver) =
471 transaction_delivery(DeliveryLimits::try_new(4, one * 3).unwrap()).unwrap();
472 delivery.event_tx.try_send(sample_event(1)).unwrap();
473 delivery.event_tx.try_send(sample_event(2)).unwrap();
474 assert!(delivery.event_tx.queued_bytes() > 0);
475 drop(receiver);
476 tokio::task::yield_now().await;
478 assert_eq!(
479 delivery.event_tx.queued_bytes(),
480 0,
481 "receiver drop must release all queued byte permits"
482 );
483 let err = delivery.event_tx.try_send(sample_event(3)).unwrap_err();
485 assert_eq!(err, EventEnqueueError::Closed);
486 }
487
488 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
490 async fn event_byte_capacity_concurrent_send_recv_no_leak() {
491 let ev = sample_event(1);
492 let one = estimate_event_bytes(&ev);
493 let (delivery, mut receiver) =
494 transaction_delivery(DeliveryLimits::try_new(32, one * 8).unwrap()).unwrap();
495 let tx = delivery.event_tx.clone();
496 let sent = Arc::new(AtomicU64::new(0));
497 let mut joins = Vec::new();
498 for t in 0..4u64 {
499 let tx = tx.clone();
500 let sent = Arc::clone(&sent);
501 joins.push(tokio::spawn(async move {
502 for i in 0..32u64 {
503 let event = sample_event(t * 100 + i + 1);
504 loop {
505 match tx.try_send(event.clone()) {
506 Ok(()) => {
507 sent.fetch_add(1, Ordering::SeqCst);
508 break;
509 }
510 Err(EventEnqueueError::ByteCapacityExceeded)
511 | Err(EventEnqueueError::ItemCapacityExceeded) => {
512 tokio::task::yield_now().await;
513 }
514 Err(e) => panic!("unexpected enqueue error: {e:?}"),
515 }
516 }
517 }
518 }));
519 }
520 let drain = tokio::spawn(async move {
521 let mut n = 0u64;
522 while n < 128 {
523 if let Some(_ev) = receiver.events.recv().await {
524 n += 1;
525 } else {
526 break;
527 }
528 }
529 n
530 });
531 for j in joins {
532 j.await.unwrap();
533 }
534 drop(tx);
535 let drained = drain.await.unwrap();
536 assert_eq!(sent.load(Ordering::SeqCst), 128);
537 assert_eq!(drained, 128);
538 assert_eq!(
539 delivery.event_tx.queued_bytes(),
540 0,
541 "no leaked byte capacity after concurrent drain"
542 );
543 }
544}