#![no_std]
extern crate alloc;
#[cfg(test)]
#[path = "attacks/mod.rs"]
pub mod attacks;
use alloc::{boxed::Box, sync::Arc};
use core::{
cell::UnsafeCell,
fmt,
future::Future,
hash::{Hash, Hasher},
mem::{ManuallyDrop, MaybeUninit},
pin::Pin,
ptr,
sync::atomic::{AtomicBool, AtomicPtr, AtomicUsize, Ordering},
task::{Context, Poll, Waker},
};
#[cfg(target_pointer_width = "64")]
const BLOCK_CAP: usize = 32;
#[cfg(target_pointer_width = "32")]
const BLOCK_CAP: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SendError<T>(pub T);
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "sending on a closed channel")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TryRecvError {
Empty,
Disconnected,
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TryRecvError::Empty => write!(f, "channel is empty"),
TryRecvError::Disconnected => write!(f, "channel is disconnected"),
}
}
}
pub fn unbounded_channel<T>() -> (UnboundedSender<T>, UnboundedReceiver<T>) {
let block = Block::new(0);
let block_ptr = Box::into_raw(Box::new(block));
let shared = Arc::new(Shared {
head: AtomicPtr::new(block_ptr),
tail: AtomicPtr::new(block_ptr),
rx_waker: AtomicPtr::new(ptr::null_mut()),
waker_lock: AtomicBool::new(false),
num_senders: AtomicUsize::new(1),
num_weak_senders: AtomicUsize::new(0),
closed: AtomicBool::new(false),
});
let sender = UnboundedSender {
shared: Arc::clone(&shared),
};
let receiver = UnboundedReceiver {
shared,
recv_index: 0,
};
(sender, receiver)
}
struct Block<T> {
next: AtomicPtr<Block<T>>,
start_index: usize,
values: UnsafeCell<[MaybeUninit<ManuallyDrop<T>>; BLOCK_CAP]>,
ready_slots: AtomicUsize,
len: AtomicUsize,
}
impl<T> Block<T> {
fn new(start_index: usize) -> Self {
Self {
next: AtomicPtr::new(ptr::null_mut()),
start_index,
values: UnsafeCell::new([const { MaybeUninit::uninit() }; BLOCK_CAP]),
ready_slots: AtomicUsize::new(0),
len: AtomicUsize::new(0),
}
}
fn relative_index(&self, index: usize) -> Option<usize> {
if index >= self.start_index && index < self.start_index + BLOCK_CAP {
Some(index - self.start_index)
} else {
None
}
}
fn write(&self, relative_index: usize, value: T) -> Result<(), T> {
if relative_index >= BLOCK_CAP {
return Err(value);
}
let mask = 1 << relative_index;
let prev_ready = self.ready_slots.fetch_or(mask, Ordering::AcqRel);
if prev_ready & mask != 0 {
return Err(value);
}
unsafe {
let values = &mut *self.values.get();
values[relative_index].write(ManuallyDrop::new(value));
}
Ok(())
}
fn read(&self, relative_index: usize) -> Option<T> {
if relative_index >= BLOCK_CAP {
return None;
}
let mask = 1 << relative_index;
let prev_ready = self.ready_slots.fetch_and(!mask, Ordering::AcqRel);
if prev_ready & mask == 0 {
return None;
}
unsafe {
let values = &*self.values.get();
Some(ManuallyDrop::into_inner(
values[relative_index].assume_init_read(),
))
}
}
fn is_ready(&self, relative_index: usize) -> bool {
if relative_index >= BLOCK_CAP {
return false;
}
let mask = 1 << relative_index;
self.ready_slots.load(Ordering::Acquire) & mask != 0
}
fn ready_count(&self) -> usize {
self.ready_slots.load(Ordering::Acquire).count_ones() as usize
}
}
impl<T> Drop for Block<T> {
fn drop(&mut self) {
let ready = self.ready_slots.load(Ordering::Relaxed);
unsafe {
let values = &mut *self.values.get();
for i in 0..BLOCK_CAP {
if ready & (1 << i) != 0 {
ManuallyDrop::drop(values[i].assume_init_mut());
}
}
}
}
}
struct Shared<T> {
head: AtomicPtr<Block<T>>,
tail: AtomicPtr<Block<T>>,
rx_waker: AtomicPtr<Waker>,
waker_lock: AtomicBool,
num_senders: AtomicUsize,
num_weak_senders: AtomicUsize,
closed: AtomicBool,
}
impl<T> Shared<T> {
fn wake_receiver(&self) {
if self.rx_waker.load(Ordering::Acquire).is_null() {
return; }
while self
.waker_lock
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_err()
{
core::hint::spin_loop();
}
let waker_ptr = self.rx_waker.swap(ptr::null_mut(), Ordering::Acquire);
if !waker_ptr.is_null() {
let waker = unsafe { Box::from_raw(waker_ptr) };
self.waker_lock.store(false, Ordering::Release);
waker.wake();
} else {
self.waker_lock.store(false, Ordering::Release);
}
}
fn store_waker(&self, waker: Waker) {
while self
.waker_lock
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_err()
{
core::hint::spin_loop();
}
let new_waker_ptr = Box::into_raw(Box::new(waker));
let old_waker_ptr = self.rx_waker.swap(new_waker_ptr, Ordering::Release);
if !old_waker_ptr.is_null() {
unsafe { drop(Box::from_raw(old_waker_ptr)) };
}
self.waker_lock.store(false, Ordering::Release);
}
}
unsafe impl<T: Send> Send for Shared<T> {}
unsafe impl<T: Send> Sync for Shared<T> {}
impl<T> Drop for Shared<T> {
fn drop(&mut self) {
let waker_ptr = self.rx_waker.load(Ordering::Relaxed);
if !waker_ptr.is_null() {
unsafe { drop(Box::from_raw(waker_ptr)) };
}
let mut current = self.tail.load(Ordering::Relaxed);
while !current.is_null() {
let block = unsafe { Box::from_raw(current) };
current = block.next.load(Ordering::Relaxed);
}
}
}
pub struct UnboundedSender<T> {
shared: Arc<Shared<T>>,
}
pub struct WeakUnboundedSender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for WeakUnboundedSender<T> {
fn clone(&self) -> Self {
self.shared.num_weak_senders.fetch_add(1, Ordering::Relaxed);
Self {
shared: Arc::clone(&self.shared),
}
}
}
impl<T> Drop for WeakUnboundedSender<T> {
fn drop(&mut self) {
self.shared.num_weak_senders.fetch_sub(1, Ordering::AcqRel);
}
}
impl<T> Clone for UnboundedSender<T> {
fn clone(&self) -> Self {
self.shared.num_senders.fetch_add(1, Ordering::Relaxed);
Self {
shared: Arc::clone(&self.shared),
}
}
}
impl<T> Drop for UnboundedSender<T> {
fn drop(&mut self) {
let prev_count = self.shared.num_senders.fetch_sub(1, Ordering::AcqRel);
if prev_count == 1 {
self.shared.closed.store(true, Ordering::Release);
self.shared.wake_receiver();
}
}
}
impl<T> UnboundedSender<T> {
pub fn send(&self, mut value: T) -> Result<(), SendError<T>> {
if self.shared.closed.load(Ordering::Acquire) {
return Err(SendError(value));
}
let mut attempts = 0;
loop {
let head_ptr = self.shared.head.load(Ordering::Acquire);
let head = unsafe { &*head_ptr };
let slot_idx = head.len.fetch_add(1, Ordering::AcqRel);
if slot_idx < BLOCK_CAP {
match head.write(slot_idx, value) {
Ok(()) => {
self.shared.wake_receiver();
return Ok(());
}
Err(returned_value) => {
value = returned_value;
head.len.fetch_sub(1, Ordering::AcqRel); continue;
}
}
} else {
head.len.store(BLOCK_CAP, Ordering::Release);
}
let next_ptr = head.next.load(Ordering::Acquire);
if next_ptr.is_null() {
let new_block = Box::into_raw(Box::new(Block::new(head.start_index + BLOCK_CAP)));
match head.next.compare_exchange_weak(
ptr::null_mut(),
new_block,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
self.shared.head.store(new_block, Ordering::Release);
}
Err(_) => {
unsafe { drop(Box::from_raw(new_block)) };
}
}
} else {
self.shared
.head
.compare_exchange_weak(head_ptr, next_ptr, Ordering::AcqRel, Ordering::Acquire)
.ok();
}
attempts += 1;
if attempts > 1000 {
core::hint::spin_loop();
attempts = 0;
}
}
}
pub fn id(&self) -> usize {
Arc::as_ptr(&self.shared) as usize
}
pub fn is_closed(&self) -> bool {
self.shared.closed.load(Ordering::Acquire)
}
pub fn same_channel(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.shared, &other.shared)
}
#[must_use = "Downgrade creates a WeakSender without destroying the original non-weak sender."]
pub fn downgrade(&self) -> WeakUnboundedSender<T> {
self.shared.num_weak_senders.fetch_add(1, Ordering::Relaxed);
WeakUnboundedSender {
shared: Arc::clone(&self.shared),
}
}
pub fn strong_count(&self) -> usize {
self.shared.num_senders.load(Ordering::Acquire)
}
pub fn weak_count(&self) -> usize {
self.shared.num_weak_senders.load(Ordering::Acquire)
}
pub async fn closed(&self) {
ClosedFuture { sender: self }.await
}
}
impl<T> WeakUnboundedSender<T> {
pub fn upgrade(&self) -> Option<UnboundedSender<T>> {
let mut count = self.shared.num_senders.load(Ordering::Acquire);
loop {
if count == 0 {
return None;
}
match self.shared.num_senders.compare_exchange_weak(
count,
count + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
return Some(UnboundedSender {
shared: Arc::clone(&self.shared),
});
}
Err(actual) => count = actual,
}
}
}
pub fn strong_count(&self) -> usize {
self.shared.num_senders.load(Ordering::Acquire)
}
pub fn weak_count(&self) -> usize {
self.shared.num_weak_senders.load(Ordering::Acquire)
}
}
impl<T> PartialEq for UnboundedSender<T> {
fn eq(&self, other: &Self) -> bool {
self.id() == other.id()
}
}
impl<T> Eq for UnboundedSender<T> {}
impl<T> PartialOrd for UnboundedSender<T> {
fn partial_cmp(&self, other: &Self) -> Option<core::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl<T> Ord for UnboundedSender<T> {
fn cmp(&self, other: &Self) -> core::cmp::Ordering {
self.id().cmp(&other.id())
}
}
impl<T> Hash for UnboundedSender<T> {
fn hash<H: Hasher>(&self, state: &mut H) {
self.id().hash(state);
}
}
impl<T> PartialEq for WeakUnboundedSender<T> {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.shared, &other.shared)
}
}
impl<T> Eq for WeakUnboundedSender<T> {}
impl<T> PartialOrd for WeakUnboundedSender<T> {
fn partial_cmp(&self, other: &Self) -> Option<core::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl<T> Ord for WeakUnboundedSender<T> {
fn cmp(&self, other: &Self) -> core::cmp::Ordering {
let self_ptr = Arc::as_ptr(&self.shared) as usize;
let other_ptr = Arc::as_ptr(&other.shared) as usize;
self_ptr.cmp(&other_ptr)
}
}
impl<T> Hash for WeakUnboundedSender<T> {
fn hash<H: Hasher>(&self, state: &mut H) {
let ptr = Arc::as_ptr(&self.shared) as usize;
ptr.hash(state);
}
}
impl<T> fmt::Debug for UnboundedSender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UnboundedSender")
.field("id", &self.id())
.field("strong_count", &self.strong_count())
.field("weak_count", &self.weak_count())
.finish()
}
}
impl<T> fmt::Debug for WeakUnboundedSender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WeakUnboundedSender")
.field("strong_count", &self.strong_count())
.field("weak_count", &self.weak_count())
.finish()
}
}
pub struct UnboundedReceiver<T> {
shared: Arc<Shared<T>>,
recv_index: usize,
}
impl<T> UnboundedReceiver<T> {
pub async fn recv(&mut self) -> Option<T> {
RecvFuture { receiver: self }.await
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
loop {
let tail_ptr = self.shared.tail.load(Ordering::Acquire);
let tail = unsafe { &*tail_ptr };
if let Some(relative_idx) = tail.relative_index(self.recv_index) {
if tail.is_ready(relative_idx) {
if let Some(value) = tail.read(relative_idx) {
self.recv_index += 1;
return Ok(value);
}
}
if relative_idx == BLOCK_CAP - 1
|| tail.ready_count() == tail.len.load(Ordering::Acquire)
{
let next_ptr = tail.next.load(Ordering::Acquire);
if !next_ptr.is_null() {
if self
.shared
.tail
.compare_exchange(
tail_ptr,
next_ptr,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
unsafe { drop(Box::from_raw(tail_ptr)) };
}
continue;
}
}
} else {
let next_ptr = tail.next.load(Ordering::Acquire);
if !next_ptr.is_null() {
if self
.shared
.tail
.compare_exchange(tail_ptr, next_ptr, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
unsafe { drop(Box::from_raw(tail_ptr)) };
}
continue;
}
}
if self.shared.closed.load(Ordering::Acquire)
&& self.shared.num_senders.load(Ordering::Acquire) == 0
{
return Err(TryRecvError::Disconnected);
}
return Err(TryRecvError::Empty);
}
}
pub fn id(&self) -> usize {
Arc::as_ptr(&self.shared) as usize
}
pub fn is_closed(&self) -> bool {
self.shared.closed.load(Ordering::Acquire)
&& self.shared.num_senders.load(Ordering::Acquire) == 0
}
pub fn is_empty(&self) -> bool {
let tail_ptr = self.shared.tail.load(Ordering::Acquire);
let tail = unsafe { &*tail_ptr };
if let Some(relative_idx) = tail.relative_index(self.recv_index) {
if tail.is_ready(relative_idx) {
return false; }
}
if self.shared.closed.load(Ordering::Acquire)
&& self.shared.num_senders.load(Ordering::Acquire) == 0
{
return true; }
true
}
pub fn close(&mut self) {
self.shared.closed.store(true, Ordering::Release);
}
pub fn sender_strong_count(&self) -> usize {
self.shared.num_senders.load(Ordering::Acquire)
}
pub fn sender_weak_count(&self) -> usize {
self.shared.num_weak_senders.load(Ordering::Acquire)
}
pub fn len(&self) -> usize {
let mut count = 0;
let mut current = self.shared.tail.load(Ordering::Acquire);
while !current.is_null() {
let block = unsafe { &*current };
count += block.ready_count();
current = block.next.load(Ordering::Acquire);
}
count
}
fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
match self.try_recv() {
Ok(value) => Poll::Ready(Some(value)),
Err(TryRecvError::Disconnected) => Poll::Ready(None),
Err(TryRecvError::Empty) => {
self.shared.store_waker(cx.waker().clone());
Poll::Pending
}
}
}
}
impl<T> Drop for UnboundedReceiver<T> {
fn drop(&mut self) {
self.shared.closed.store(true, Ordering::Release);
}
}
struct RecvFuture<'a, T> {
receiver: &'a mut UnboundedReceiver<T>,
}
impl<'a, T> Future for RecvFuture<'a, T> {
type Output = Option<T>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.receiver.poll_recv(cx)
}
}
struct ClosedFuture<'a, T> {
sender: &'a UnboundedSender<T>,
}
impl<'a, T> Future for ClosedFuture<'a, T> {
type Output = ();
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
if self.sender.is_closed() {
Poll::Ready(())
} else {
Poll::Pending
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::{vec, vec::Vec};
#[test]
fn test_basic_send_recv() {
let (tx, mut rx) = unbounded_channel::<i32>();
tx.send(1).unwrap();
tx.send(2).unwrap();
tx.send(3).unwrap();
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.try_recv().unwrap(), 2);
assert_eq!(rx.try_recv().unwrap(), 3);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
}
#[test]
fn test_channel_id() {
let (tx1, rx1) = unbounded_channel::<i32>();
let (tx2, rx2) = unbounded_channel::<i32>();
assert_eq!(tx1.id(), rx1.id());
assert_ne!(tx1.id(), tx2.id());
assert_ne!(rx1.id(), rx2.id());
}
#[test]
fn test_clone_sender() {
let (tx, mut rx) = unbounded_channel::<i32>();
let tx2 = tx.clone();
tx.send(1).unwrap();
tx2.send(2).unwrap();
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.try_recv().unwrap(), 2);
}
#[test]
fn test_large_number_of_messages() {
let (tx, mut rx) = unbounded_channel::<usize>();
for i in 0..100 {
tx.send(i).unwrap();
}
for i in 0..100 {
assert_eq!(rx.try_recv().unwrap(), i);
}
}
#[test]
fn test_drop_sender_closes_channel() {
let (tx, mut rx) = unbounded_channel::<i32>();
tx.send(42).unwrap();
drop(tx);
assert_eq!(rx.try_recv().unwrap(), 42);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
}
#[test]
fn test_same_channel() {
let (tx1, _rx) = unbounded_channel::<i32>();
let tx2 = tx1.clone();
let (tx3, _rx2) = unbounded_channel::<i32>();
assert!(tx1.same_channel(&tx2));
assert!(!tx1.same_channel(&tx3));
}
#[test]
fn test_stress_many_messages() {
let (tx, mut rx) = unbounded_channel::<usize>();
const NUM_MESSAGES: usize = 10_000;
for i in 0..NUM_MESSAGES {
tx.send(i).unwrap();
}
for i in 0..NUM_MESSAGES {
assert_eq!(rx.try_recv().unwrap(), i);
}
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
}
#[test]
fn test_send_recv_interleaved() {
let (tx, mut rx) = unbounded_channel::<i32>();
let mut expected_recv = 0;
for i in 0..100 {
tx.send(i).unwrap();
if i % 2 == 0 {
assert_eq!(rx.try_recv().unwrap(), expected_recv);
expected_recv += 1;
}
}
while let Ok(value) = rx.try_recv() {
assert_eq!(value, expected_recv);
expected_recv += 1;
}
assert_eq!(expected_recv, 100); }
#[test]
fn test_drop_receiver_while_sending() {
let (tx, rx) = unbounded_channel::<i32>();
tx.send(1).unwrap();
tx.send(2).unwrap();
drop(rx);
assert!(matches!(tx.send(3), Err(SendError(3))));
assert!(tx.is_closed());
}
#[test]
fn test_multiple_sender_drops() {
let (tx, mut rx) = unbounded_channel::<i32>();
let tx2 = tx.clone();
let tx3 = tx.clone();
tx.send(1).unwrap();
tx2.send(2).unwrap();
tx3.send(3).unwrap();
drop(tx);
assert!(!rx.is_closed());
drop(tx2);
assert!(!rx.is_closed());
drop(tx3);
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.try_recv().unwrap(), 2);
assert_eq!(rx.try_recv().unwrap(), 3);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
assert!(rx.is_closed());
}
#[test]
fn test_zero_sized_types() {
#[derive(Debug, PartialEq)]
struct ZeroSized;
let (tx, mut rx) = unbounded_channel::<ZeroSized>();
tx.send(ZeroSized).unwrap();
tx.send(ZeroSized).unwrap();
assert_eq!(rx.try_recv().unwrap(), ZeroSized);
assert_eq!(rx.try_recv().unwrap(), ZeroSized);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
}
#[test]
fn test_large_types() {
#[derive(Debug, PartialEq)]
struct LargeType([u8; 1024]);
let (tx, mut rx) = unbounded_channel::<LargeType>();
let large_value = LargeType([42; 1024]);
tx.send(large_value).unwrap();
let received = rx.try_recv().unwrap();
assert_eq!(received.0.len(), 1024);
for &byte in &received.0 {
assert_eq!(byte, 42, "Large type data corruption detected");
}
let large_value2 = LargeType([123; 1024]);
let large_value3 = LargeType([255; 1024]);
tx.send(large_value2).unwrap();
tx.send(large_value3).unwrap();
let received2 = rx.try_recv().unwrap();
let received3 = rx.try_recv().unwrap();
for &byte in &received2.0 {
assert_eq!(byte, 123, "Second large message corrupted");
}
for &byte in &received3.0 {
assert_eq!(byte, 255, "Third large message corrupted");
}
}
#[test]
fn test_unwind_safety_basic() {
#[derive(Debug)]
struct ConditionalPanic(bool);
impl Drop for ConditionalPanic {
fn drop(&mut self) {
if self.0 {
}
}
}
let (tx, mut rx) = unbounded_channel::<ConditionalPanic>();
tx.send(ConditionalPanic(false)).unwrap(); tx.send(ConditionalPanic(true)).unwrap(); tx.send(ConditionalPanic(false)).unwrap();
assert_eq!(rx.try_recv().unwrap().0, false);
assert_eq!(rx.try_recv().unwrap().0, true);
assert_eq!(rx.try_recv().unwrap().0, false);
tx.send(ConditionalPanic(false)).unwrap();
assert_eq!(rx.try_recv().unwrap().0, false);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
assert!(!rx.is_closed());
}
#[test]
fn test_block_boundary_conditions() {
let (tx, mut rx) = unbounded_channel::<usize>();
for i in 0..BLOCK_CAP {
tx.send(i).unwrap();
}
tx.send(BLOCK_CAP).unwrap();
for i in 0..=BLOCK_CAP {
assert_eq!(rx.try_recv().unwrap(), i);
}
for i in 0..(BLOCK_CAP * 3) {
tx.send(i).unwrap();
}
for i in 0..(BLOCK_CAP * 3) {
assert_eq!(rx.try_recv().unwrap(), i);
}
}
#[test]
fn test_receiver_state_consistency() {
let (tx, mut rx) = unbounded_channel::<i32>();
assert!(rx.is_empty());
assert!(!rx.is_closed());
tx.send(42).unwrap();
assert!(!rx.is_empty());
assert!(!rx.is_closed());
assert_eq!(rx.try_recv().unwrap(), 42);
assert!(rx.is_empty());
assert!(!rx.is_closed());
drop(tx);
assert!(rx.is_empty());
assert!(rx.is_closed());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
}
#[test]
fn test_manual_close() {
let (tx, mut rx) = unbounded_channel::<i32>();
tx.send(1).unwrap();
tx.send(2).unwrap();
rx.close();
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.try_recv().unwrap(), 2);
assert!(tx.is_closed());
assert!(matches!(tx.send(3), Err(SendError(3))));
}
#[test]
fn test_channel_id_consistency() {
let (tx, rx) = unbounded_channel::<i32>();
assert_eq!(tx.id(), rx.id());
let (tx2, rx2) = unbounded_channel::<i32>();
assert_eq!(tx2.id(), rx2.id());
let tx_clone = tx.clone();
assert_eq!(tx.id(), tx_clone.id());
assert!(tx.same_channel(&tx_clone));
assert!(!tx.same_channel(&tx2));
}
#[test]
fn test_drop_semantics() {
use alloc::rc::Rc;
let drop_count = Rc::new(core::cell::RefCell::new(0));
#[derive(Debug)]
struct DropCounter(Rc<core::cell::RefCell<i32>>);
impl Drop for DropCounter {
fn drop(&mut self) {
*self.0.borrow_mut() += 1;
}
}
let (tx, mut rx) = unbounded_channel::<DropCounter>();
tx.send(DropCounter(drop_count.clone())).unwrap();
tx.send(DropCounter(drop_count.clone())).unwrap();
tx.send(DropCounter(drop_count.clone())).unwrap();
assert_eq!(*drop_count.borrow(), 0);
let _value1 = rx.try_recv().unwrap();
assert_eq!(*drop_count.borrow(), 0);
drop(_value1);
assert_eq!(*drop_count.borrow(), 1);
drop(tx);
drop(rx);
assert_eq!(*drop_count.borrow(), 3); }
#[test]
fn test_memory_safety_after_close() {
let (tx, mut rx) = unbounded_channel::<Vec<u8>>();
tx.send(vec![1, 2, 3]).unwrap();
tx.send(vec![4, 5, 6]).unwrap();
rx.close();
assert!(matches!(tx.send(vec![7, 8, 9]), Err(_)));
assert_eq!(rx.try_recv().unwrap(), vec![1, 2, 3]);
assert_eq!(rx.try_recv().unwrap(), vec![4, 5, 6]);
}
#[test]
fn test_ordering_guarantees() {
let (tx, mut rx) = unbounded_channel::<usize>();
for i in 0..1000 {
tx.send(i).unwrap();
}
for i in 0..1000 {
assert_eq!(rx.try_recv().unwrap(), i);
}
}
#[test]
fn test_empty_channel_operations() {
let (tx, mut rx) = unbounded_channel::<i32>();
assert!(rx.is_empty());
assert!(!rx.is_closed());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
drop(tx);
assert!(rx.is_empty());
assert!(rx.is_closed());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
}
#[test]
fn test_channel_reuse_after_empty() {
let (tx, mut rx) = unbounded_channel::<i32>();
for round in 0..10 {
for i in 0..10 {
tx.send(round * 10 + i).unwrap();
}
for i in 0..10 {
assert_eq!(rx.try_recv().unwrap(), round * 10 + i);
}
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
}
}
#[test]
fn test_mixed_operation_patterns() {
let (tx, mut rx) = unbounded_channel::<usize>();
let mut next_send = 0;
let mut next_recv = 0;
for _ in 0..100 {
let send_count = (next_send % 5) + 1;
for _ in 0..send_count {
tx.send(next_send).unwrap();
next_send += 1;
}
let recv_count = (next_recv % 3) + 1;
for _ in 0..recv_count {
if let Ok(value) = rx.try_recv() {
assert_eq!(value, next_recv);
next_recv += 1;
} else {
break;
}
}
}
while let Ok(value) = rx.try_recv() {
assert_eq!(value, next_recv);
next_recv += 1;
}
assert_eq!(next_send, next_recv);
}
#[test]
fn test_weak_sender_basic() {
let (tx, mut rx) = unbounded_channel::<i32>();
let weak_tx = tx.downgrade();
let upgraded_tx = weak_tx.upgrade().unwrap();
upgraded_tx.send(42).unwrap();
assert_eq!(rx.try_recv().unwrap(), 42);
drop(tx);
upgraded_tx.send(43).unwrap();
assert_eq!(rx.try_recv().unwrap(), 43);
drop(upgraded_tx);
assert!(weak_tx.upgrade().is_none());
}
#[test]
fn test_weak_sender_upgrade_failure() {
let (tx, _rx) = unbounded_channel::<i32>();
let weak_tx = tx.downgrade();
drop(tx);
assert!(weak_tx.upgrade().is_none());
}
#[test]
fn test_weak_sender_counts() {
let (tx, rx) = unbounded_channel::<i32>();
assert_eq!(tx.strong_count(), 1);
assert_eq!(tx.weak_count(), 0);
assert_eq!(rx.sender_strong_count(), 1);
assert_eq!(rx.sender_weak_count(), 0);
let weak_tx = tx.downgrade();
assert_eq!(tx.strong_count(), 1);
assert_eq!(tx.weak_count(), 1);
assert_eq!(weak_tx.strong_count(), 1);
assert_eq!(weak_tx.weak_count(), 1);
assert_eq!(rx.sender_strong_count(), 1);
assert_eq!(rx.sender_weak_count(), 1);
let tx2 = tx.clone();
assert_eq!(tx.strong_count(), 2);
assert_eq!(tx.weak_count(), 1);
assert_eq!(tx2.strong_count(), 2);
assert_eq!(weak_tx.strong_count(), 2);
assert_eq!(weak_tx.weak_count(), 1);
assert_eq!(rx.sender_strong_count(), 2);
assert_eq!(rx.sender_weak_count(), 1);
let weak_tx2 = weak_tx.clone();
assert_eq!(tx.strong_count(), 2);
assert_eq!(tx.weak_count(), 2);
assert_eq!(weak_tx.weak_count(), 2);
assert_eq!(weak_tx2.weak_count(), 2);
assert_eq!(rx.sender_strong_count(), 2);
assert_eq!(rx.sender_weak_count(), 2);
drop(weak_tx);
assert_eq!(tx.weak_count(), 1);
assert_eq!(weak_tx2.weak_count(), 1);
assert_eq!(rx.sender_weak_count(), 1);
drop(tx);
assert_eq!(tx2.strong_count(), 1);
assert_eq!(weak_tx2.strong_count(), 1);
assert_eq!(rx.sender_strong_count(), 1);
drop(tx2);
assert_eq!(weak_tx2.strong_count(), 0);
assert_eq!(weak_tx2.weak_count(), 1);
assert_eq!(rx.sender_strong_count(), 0);
assert_eq!(rx.sender_weak_count(), 1);
assert!(weak_tx2.upgrade().is_none());
}
#[test]
fn test_weak_sender_channel_close() {
let (tx, rx) = unbounded_channel::<i32>();
let weak_tx = tx.downgrade();
drop(tx);
assert!(rx.is_closed());
assert!(weak_tx.upgrade().is_none());
}
#[test]
fn test_sender_ordering_and_equality() {
let (tx1, _rx1) = unbounded_channel::<i32>();
let (tx2, _rx2) = unbounded_channel::<i32>();
let tx1_clone = tx1.clone();
let weak_tx1 = tx1.downgrade();
let weak_tx2 = tx2.downgrade();
assert_eq!(tx1, tx1_clone);
assert_ne!(tx1, tx2);
assert_eq!(weak_tx1, weak_tx1.clone());
assert_ne!(weak_tx1, weak_tx2);
let ordering1 = tx1.cmp(&tx2);
let ordering2 = tx1.cmp(&tx2);
assert_eq!(ordering1, ordering2);
use alloc::collections::BTreeSet;
let mut set = BTreeSet::new();
set.insert(tx1.clone());
set.insert(tx1_clone.clone());
assert_eq!(set.len(), 1);
set.insert(tx2.clone());
assert_eq!(set.len(), 2); }
#[test]
fn test_weak_sender_multiple_upgrades() {
let (tx, mut rx) = unbounded_channel::<i32>();
let weak_tx = tx.downgrade();
let upgraded1 = weak_tx.upgrade().unwrap();
let upgraded2 = weak_tx.upgrade().unwrap();
upgraded1.send(1).unwrap();
upgraded2.send(2).unwrap();
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.try_recv().unwrap(), 2);
drop(tx);
drop(upgraded1);
let upgraded3 = weak_tx.upgrade().unwrap();
upgraded3.send(3).unwrap();
assert_eq!(rx.try_recv().unwrap(), 3);
drop(upgraded2);
drop(upgraded3);
assert!(weak_tx.upgrade().is_none());
}
#[test]
fn test_sender_hash_collections() {
use alloc::collections::BTreeSet;
let (tx1, _rx1) = unbounded_channel::<i32>();
let (tx2, _rx2) = unbounded_channel::<i32>();
let tx1_clone = tx1.clone();
let mut set = BTreeSet::new();
set.insert(tx1.clone());
assert_eq!(set.len(), 1);
set.insert(tx1_clone);
assert_eq!(set.len(), 1);
set.insert(tx2);
assert_eq!(set.len(), 2);
let weak_tx1 = tx1.downgrade();
let weak_tx1_clone = weak_tx1.clone();
let mut weak_set = BTreeSet::new();
weak_set.insert(weak_tx1);
weak_set.insert(weak_tx1_clone); assert_eq!(weak_set.len(), 1);
}
#[test]
fn test_len_method() {
let (tx, mut rx) = unbounded_channel::<i32>();
assert_eq!(rx.len(), 0);
assert!(rx.is_empty());
tx.send(1).unwrap();
assert_eq!(rx.len(), 1);
assert!(!rx.is_empty());
tx.send(2).unwrap();
tx.send(3).unwrap();
assert_eq!(rx.len(), 3);
assert!(!rx.is_empty());
assert_eq!(rx.try_recv().unwrap(), 1);
assert_eq!(rx.len(), 2);
assert!(!rx.is_empty());
assert_eq!(rx.try_recv().unwrap(), 2);
assert_eq!(rx.len(), 1);
assert!(!rx.is_empty());
assert_eq!(rx.try_recv().unwrap(), 3);
assert_eq!(rx.len(), 0);
assert!(rx.is_empty());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
}
#[test]
fn test_tokio_drop_behavior_compatibility() {
use alloc::sync::Arc;
use core::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
struct DropCounter {
#[allow(dead_code)]
id: usize,
counter: Arc<AtomicUsize>,
}
impl Drop for DropCounter {
fn drop(&mut self) {
self.counter.fetch_add(1, Ordering::SeqCst);
}
}
const NUM_MESSAGES: usize = 100;
let our_drop_counter = Arc::new(AtomicUsize::new(0));
{
let (tx, mut rx) = unbounded_channel::<DropCounter>();
for i in 0..NUM_MESSAGES {
let msg = DropCounter {
id: i,
counter: Arc::clone(&our_drop_counter),
};
tx.send(msg).unwrap();
}
for _ in 0..10 {
let _msg = rx.try_recv().unwrap();
}
drop(rx);
drop(tx);
}
let our_dropped_count = our_drop_counter.load(Ordering::SeqCst);
assert_eq!(
our_dropped_count, NUM_MESSAGES,
"Our implementation should drop all {} messages",
NUM_MESSAGES
);
}
#[test]
fn test_tokio_exact_drop_condition() {
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
struct DropTracker {
#[allow(dead_code)]
id: usize,
counter: Arc<AtomicUsize>,
}
impl Drop for DropTracker {
fn drop(&mut self) {
self.counter.fetch_add(1, Ordering::SeqCst);
}
}
const NUM_MESSAGES: usize = 50; let drop_counter = Arc::new(AtomicUsize::new(0));
let (tx, mut rx) = unbounded_channel::<DropTracker>();
let tx_clone = tx.clone();
for i in 0..NUM_MESSAGES {
let msg = DropTracker {
id: i,
counter: Arc::clone(&drop_counter),
};
tx.send(msg).unwrap();
}
let received_before_drop = 10;
for _ in 0..received_before_drop {
let _msg = rx.try_recv().unwrap();
}
drop(rx);
let mut send_errors = 0;
let mut failed_messages = Vec::new();
for i in NUM_MESSAGES..NUM_MESSAGES + 10 {
let msg = DropTracker {
id: i,
counter: Arc::clone(&drop_counter),
};
match tx_clone.send(msg) {
Ok(_) => {
panic!("Send succeeded after receiver drop - this violates Tokio behavior");
}
Err(send_error) => {
send_errors += 1;
failed_messages.push(send_error.0); }
}
}
drop(tx);
drop(tx_clone);
drop(failed_messages);
let final_drop_count = drop_counter.load(Ordering::SeqCst);
let expected_total_drops = NUM_MESSAGES + send_errors;
assert_eq!(
final_drop_count, expected_total_drops,
"Expected {} drops (original {} + failed sends {}), got {}",
expected_total_drops, NUM_MESSAGES, send_errors, final_drop_count
);
assert!(
send_errors > 0,
"Send attempts after receiver drop should fail"
);
}
#[test]
fn test_tokio_api_compatibility() {
let (tx, mut rx) = unbounded_channel::<i32>();
assert!(!tx.is_closed());
assert!(!rx.is_closed());
assert!(rx.is_empty());
assert_eq!(rx.len(), 0);
assert_eq!(tx.strong_count(), 1);
assert_eq!(tx.weak_count(), 0);
assert_eq!(rx.sender_strong_count(), 1);
assert_eq!(rx.sender_weak_count(), 0);
assert!(tx.same_channel(&tx));
assert_eq!(tx.id(), rx.id());
let _weak_tx = tx.downgrade();
assert_eq!(tx.weak_count(), 1);
assert_eq!(rx.sender_weak_count(), 1);
tx.send(42).unwrap();
assert_eq!(rx.len(), 1);
assert!(!rx.is_empty());
assert_eq!(rx.try_recv().unwrap(), 42);
assert_eq!(rx.len(), 0);
assert!(rx.is_empty());
}
#[test]
fn test_sender_drop_behavior_comprehensive() {
use alloc::sync::Arc;
use core::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
struct DropTracker {
#[allow(dead_code)]
id: usize,
counter: Arc<AtomicUsize>,
}
impl Drop for DropTracker {
fn drop(&mut self) {
self.counter.fetch_add(1, Ordering::SeqCst);
}
}
{
let drop_counter = Arc::new(AtomicUsize::new(0));
const NUM_MESSAGES: usize = 50;
let (tx, mut rx) = unbounded_channel::<DropTracker>();
let tx2 = tx.clone();
let tx3 = tx.clone();
for i in 0..NUM_MESSAGES {
let msg = DropTracker {
id: i,
counter: Arc::clone(&drop_counter),
};
match i % 3 {
0 => tx.send(msg).unwrap(),
1 => tx2.send(msg).unwrap(),
2 => tx3.send(msg).unwrap(),
_ => unreachable!(),
}
}
let received_count = 15;
for _ in 0..received_count {
let _msg = rx.try_recv().unwrap();
}
assert_eq!(rx.len(), NUM_MESSAGES - received_count);
assert!(!rx.is_closed());
drop(tx);
drop(tx2);
drop(tx3);
assert!(rx.is_closed());
let mut remaining_received = 0;
while let Ok(_msg) = rx.try_recv() {
remaining_received += 1;
}
assert_eq!(remaining_received, NUM_MESSAGES - received_count);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
drop(rx);
assert_eq!(
drop_counter.load(Ordering::SeqCst),
NUM_MESSAGES,
"All messages should be dropped when senders and receiver are dropped"
);
}
{
let drop_counter = Arc::new(AtomicUsize::new(0));
const NUM_MESSAGES: usize = 30;
let (tx, mut rx) = unbounded_channel::<DropTracker>();
let tx2 = tx.clone();
for i in 0..NUM_MESSAGES {
let msg = DropTracker {
id: i + 100, counter: Arc::clone(&drop_counter),
};
if i % 2 == 0 {
tx.send(msg).unwrap();
} else {
tx2.send(msg).unwrap();
}
let _received = rx.try_recv().unwrap();
}
assert!(rx.is_empty());
assert!(!rx.is_closed());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Empty)));
drop(tx);
drop(tx2);
assert!(rx.is_empty());
assert!(rx.is_closed());
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
drop(rx);
assert_eq!(
drop_counter.load(Ordering::SeqCst),
NUM_MESSAGES,
"All messages should be dropped when received"
);
}
{
let drop_counter = Arc::new(AtomicUsize::new(0));
const NUM_MESSAGES: usize = 40;
let (tx, mut rx) = unbounded_channel::<DropTracker>();
let tx2 = tx.clone();
let weak_tx = tx.downgrade();
for i in 0..NUM_MESSAGES / 2 {
let msg = DropTracker {
id: i + 200, counter: Arc::clone(&drop_counter),
};
tx.send(msg).unwrap();
}
drop(tx);
assert!(!rx.is_closed());
for i in NUM_MESSAGES / 2..NUM_MESSAGES {
let msg = DropTracker {
id: i + 200,
counter: Arc::clone(&drop_counter),
};
tx2.send(msg).unwrap();
}
let upgraded = weak_tx.upgrade();
assert!(upgraded.is_some());
drop(upgraded);
drop(tx2);
let mut received_all = 0;
while let Ok(_msg) = rx.try_recv() {
received_all += 1;
}
assert_eq!(received_all, NUM_MESSAGES);
assert!(matches!(rx.try_recv(), Err(TryRecvError::Disconnected)));
assert!(rx.is_closed());
assert!(weak_tx.upgrade().is_none());
drop(rx);
drop(weak_tx);
assert_eq!(
drop_counter.load(Ordering::SeqCst),
NUM_MESSAGES,
"All messages should be dropped with gradual sender drop"
);
}
}
}