use std::sync::Arc;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
pub struct Envelope<T> {
value: T,
permit: OwnedSemaphorePermit,
}
impl<T> Envelope<T> {
pub fn into_inner(self) -> T {
self.value
}
pub fn split(self) -> (T, BudgetPermit) {
(
self.value,
BudgetPermit {
_permit: self.permit,
},
)
}
}
impl<T> std::ops::Deref for Envelope<T> {
type Target = T;
fn deref(&self) -> &T {
&self.value
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for Envelope<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("Envelope").field(&self.value).finish()
}
}
impl<T: PartialEq> PartialEq for Envelope<T> {
fn eq(&self, other: &Self) -> bool {
self.value == other.value
}
}
impl<T: Eq> Eq for Envelope<T> {}
pub struct BudgetPermit {
_permit: OwnedSemaphorePermit,
}
#[derive(Debug)]
pub enum TrySendError<T> {
BudgetExceeded(T),
Channel(mpsc::error::TrySendError<T>),
}
#[derive(Debug)]
pub struct SendError<T>(pub T);
pub struct ByteBoundedSender<T> {
inner: mpsc::Sender<Envelope<T>>,
budget: Arc<Semaphore>,
max_queued_bytes: usize,
}
impl<T> Clone for ByteBoundedSender<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
budget: Arc::clone(&self.budget),
max_queued_bytes: self.max_queued_bytes,
}
}
}
fn permits_for(payload_len: usize) -> u32 {
u32::try_from(payload_len.max(1)).unwrap_or(u32::MAX)
}
impl<T> ByteBoundedSender<T> {
pub fn try_send(&self, value: T, payload_len: usize) -> Result<(), TrySendError<T>> {
if payload_len.max(1) > self.max_queued_bytes {
return Err(TrySendError::BudgetExceeded(value));
}
let permit = match Arc::clone(&self.budget).try_acquire_many_owned(permits_for(payload_len))
{
Ok(permit) => permit,
Err(_) => return Err(TrySendError::BudgetExceeded(value)),
};
self.inner
.try_send(Envelope { value, permit })
.map_err(|err| match err {
mpsc::error::TrySendError::Full(envelope) => {
TrySendError::Channel(mpsc::error::TrySendError::Full(envelope.value))
}
mpsc::error::TrySendError::Closed(envelope) => {
TrySendError::Channel(mpsc::error::TrySendError::Closed(envelope.value))
}
})
}
pub async fn send(&self, value: T, payload_len: usize) -> Result<(), SendError<T>> {
if payload_len.max(1) > self.max_queued_bytes {
return Err(SendError(value));
}
let permit = match Arc::clone(&self.budget)
.acquire_many_owned(permits_for(payload_len))
.await
{
Ok(permit) => permit,
Err(_) => return Err(SendError(value)),
};
self.inner
.send(Envelope { value, permit })
.await
.map_err(|err| SendError(err.0.value))
}
pub fn capacity(&self) -> usize {
self.inner.capacity()
}
pub fn max_capacity(&self) -> usize {
self.inner.max_capacity()
}
}
pub fn byte_bounded_channel<T>(
item_capacity: usize,
max_queued_bytes: usize,
) -> (ByteBoundedSender<T>, mpsc::Receiver<Envelope<T>>) {
let (inner, rx) = mpsc::channel(item_capacity);
(
ByteBoundedSender {
inner,
budget: Arc::new(Semaphore::new(max_queued_bytes)),
max_queued_bytes,
},
rx,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_rejects_send_once_item_capacity_is_full() {
let (tx, mut rx) = byte_bounded_channel::<&str>(2, 1024);
tx.try_send("a", 1).unwrap();
tx.try_send("b", 1).unwrap();
assert!(matches!(
tx.try_send("c", 1),
Err(TrySendError::Channel(mpsc::error::TrySendError::Full(_)))
));
rx.recv().await.unwrap();
tx.try_send("c", 1).expect("capacity freed after a receive");
}
#[tokio::test]
async fn test_rejects_send_once_byte_budget_is_exceeded() {
let (tx, _rx) = byte_bounded_channel::<&str>(100, 10);
assert!(matches!(
tx.try_send("this string is way over budget", 31),
Err(TrySendError::BudgetExceeded(_))
));
tx.try_send("ok", 2).expect("small payload fits in budget");
}
#[tokio::test]
async fn test_byte_budget_is_released_when_envelope_is_received() {
let (tx, mut rx) = byte_bounded_channel::<&str>(100, 10);
tx.try_send("12345", 5).unwrap();
assert!(matches!(
tx.try_send("123456", 6),
Err(TrySendError::BudgetExceeded(_))
));
let received = rx.recv().await.unwrap();
assert_eq!(received.into_inner(), "12345");
tx.try_send("123456", 6)
.expect("budget freed after the first item was received and dropped");
}
#[tokio::test]
async fn test_try_send_rejects_payload_larger_than_total_budget() {
let (tx, _rx) = byte_bounded_channel::<&str>(100, 10);
assert!(matches!(
tx.try_send("too big", 11),
Err(TrySendError::BudgetExceeded(_))
));
}
#[tokio::test]
async fn test_send_rejects_payload_larger_than_total_budget_instead_of_hanging() {
let (tx, _rx) = byte_bounded_channel::<&str>(100, 10);
let result = tokio::time::timeout(
std::time::Duration::from_millis(200),
tx.send("too big", 11),
)
.await
.expect("send must reject an unsatisfiable payload immediately, not hang");
assert!(result.is_err());
tx.try_send("ok", 5)
.expect("channel must remain usable after rejecting an oversized send");
}
#[tokio::test]
async fn test_split_defers_budget_release_until_permit_is_dropped() {
let (tx, mut rx) = byte_bounded_channel::<&str>(100, 10);
tx.try_send("12345", 5).unwrap();
let envelope = rx.recv().await.unwrap();
let (value, permit) = envelope.split();
assert_eq!(value, "12345");
assert!(matches!(
tx.try_send("123456", 6),
Err(TrySendError::BudgetExceeded(_))
));
drop(permit);
tx.try_send("123456", 6)
.expect("budget freed once the split-off permit is dropped");
}
#[tokio::test]
async fn test_concurrent_try_send_never_admits_past_byte_budget() {
let (tx, mut rx) = byte_bounded_channel::<usize>(1000, 100);
let payload_len = 10;
let mut handles = Vec::new();
for i in 0..20 {
let tx = tx.clone();
handles.push(tokio::spawn(
async move { tx.try_send(i, payload_len).is_ok() },
));
}
let mut admitted = 0;
for handle in handles {
if handle.await.expect("task panicked") {
admitted += 1;
}
}
assert!(
admitted <= 10,
"a 100-byte budget at 10 bytes/item must never admit more than 10 \
concurrent items, got {admitted}"
);
let mut drained = 0;
while rx.try_recv().is_ok() {
drained += 1;
}
assert_eq!(
drained, admitted,
"every admitted item must be receivable exactly once"
);
}
#[tokio::test]
async fn test_send_waits_for_budget_then_succeeds() {
let (tx, mut rx) = byte_bounded_channel::<&str>(100, 5);
tx.try_send("abcde", 5).unwrap();
let tx2 = tx.clone();
let waiter = tokio::spawn(async move { tx2.send("fghij", 5).await });
tokio::task::yield_now().await;
assert!(!waiter.is_finished());
rx.recv().await.unwrap();
waiter
.await
.expect("task panicked")
.expect("send should succeed once budget frees up");
}
}