1use super::wait_queue::WaitQueue;
6use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
7use std::sync::{Arc, RwLock};
8use std::task::{Context, Poll};
9
10#[repr(align(64))]
11struct Shared<T> {
12 value: RwLock<T>,
13 version: AtomicU64,
17 sender_count: AtomicUsize,
18 receiver_count: AtomicUsize,
19 wait: WaitQueue,
20}
21
22#[must_use]
24#[inline]
25pub fn channel<T>(init: T) -> (Sender<T>, Receiver<T>) {
26 let shared = Arc::new(Shared {
27 value: RwLock::new(init),
28 version: AtomicU64::new(0),
29 sender_count: AtomicUsize::new(1),
30 receiver_count: AtomicUsize::new(1),
31 wait: WaitQueue::new(),
32 });
33 let receiver = Receiver {
34 shared: shared.clone(),
35 seen_version: 0,
36 };
37 (Sender { shared }, receiver)
38}
39
40#[repr(align(64))]
42pub struct Sender<T> {
43 shared: Arc<Shared<T>>,
44}
45
46impl<T> Clone for Sender<T> {
47 #[inline(always)]
48 fn clone(&self) -> Self {
49 self.shared.sender_count.fetch_add(1, Ordering::AcqRel);
50 Self {
51 shared: self.shared.clone(),
52 }
53 }
54}
55
56impl<T> Drop for Sender<T> {
57 #[inline(always)]
58 fn drop(&mut self) {
59 if self.shared.sender_count.fetch_sub(1, Ordering::AcqRel) == 1 {
60 self.shared.wait.wake_all();
70 }
71 }
72}
73
74impl<T> Sender<T> {
75 #[inline(always)]
77 pub fn send(&self, value: T) {
78 *self
79 .shared
80 .value
81 .write()
82 .unwrap_or_else(std::sync::PoisonError::into_inner) = value;
83 self.shared.version.fetch_add(1, Ordering::Release);
84 self.shared.wait.wake_all();
85 }
86
87 #[inline(always)]
89 pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
90 self.shared
91 .value
92 .read()
93 .unwrap_or_else(std::sync::PoisonError::into_inner)
94 }
95
96 #[must_use]
98 #[inline(always)]
99 pub fn is_closed(&self) -> bool {
100 self.shared.receiver_count.load(Ordering::Acquire) == 0
101 }
102}
103
104#[repr(align(64))]
110pub struct Receiver<T> {
111 shared: Arc<Shared<T>>,
112 seen_version: u64,
113}
114
115impl<T> Clone for Receiver<T> {
116 #[inline(always)]
117 fn clone(&self) -> Self {
118 self.shared.receiver_count.fetch_add(1, Ordering::AcqRel);
119 Self {
120 shared: self.shared.clone(),
121 seen_version: self.seen_version,
122 }
123 }
124}
125
126impl<T> Drop for Receiver<T> {
127 #[inline(always)]
128 fn drop(&mut self) {
129 self.shared.receiver_count.fetch_sub(1, Ordering::AcqRel);
130 }
131}
132
133impl<T: Send + Sync> Receiver<T> {
134 #[inline(always)]
142 pub async fn changed(&mut self) -> Result<(), RecvError> {
143 std::future::poll_fn(|cx| self.poll_changed(cx)).await
144 }
145}
146
147impl<T> Receiver<T> {
148 #[inline(always)]
152 pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
153 self.shared
154 .value
155 .read()
156 .unwrap_or_else(std::sync::PoisonError::into_inner)
157 }
158
159 #[inline(always)]
163 pub fn borrow_and_update(&mut self) -> std::sync::RwLockReadGuard<'_, T> {
164 self.seen_version = self.shared.version.load(Ordering::Acquire);
165 self.borrow()
166 }
167
168 #[inline]
169 fn poll_changed(&mut self, cx: &Context<'_>) -> Poll<Result<(), RecvError>> {
170 if self.try_observe_change() {
171 return Poll::Ready(Ok(()));
172 }
173 if self.is_closed() {
174 return Poll::Ready(if self.try_observe_change() {
179 Ok(())
180 } else {
181 Err(RecvError)
182 });
183 }
184 let token = self.shared.wait.register(cx.waker());
185 if self.try_observe_change() {
186 self.shared.wait.cancel(token);
187 return Poll::Ready(Ok(()));
188 }
189 if self.is_closed() {
190 let result = if self.try_observe_change() {
191 Ok(())
192 } else {
193 Err(RecvError)
194 };
195 self.shared.wait.cancel(token);
196 return Poll::Ready(result);
197 }
198 Poll::Pending
199 }
200
201 #[inline(always)]
202 fn try_observe_change(&mut self) -> bool {
203 let current = self.shared.version.load(Ordering::Acquire);
204 if current == self.seen_version {
205 false
206 } else {
207 self.seen_version = current;
208 true
209 }
210 }
211
212 #[inline(always)]
213 fn is_closed(&self) -> bool {
214 self.shared.sender_count.load(Ordering::Acquire) == 0
215 }
216}
217
218#[derive(Debug, Clone, Copy, PartialEq, Eq)]
221#[repr(align(64))]
222pub struct RecvError;
223
224impl std::fmt::Display for RecvError {
225 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
226 f.write_str("channel closed: every sender dropped")
227 }
228}
229
230impl std::error::Error for RecvError {}