1use std::pin::Pin;
11use std::sync::{
12 Arc,
13 atomic::{AtomicBool, Ordering},
14};
15use std::task::{Context, Poll};
16use std::thread;
17use std::time::Duration;
18
19use futures::Stream;
20use futures::channel::mpsc::{UnboundedReceiver, unbounded};
21pub use net_lattice_core::{Error, Result};
22use net_lattice_platform::{EventReceiver, TokioEventReceiver};
23
24pub struct EventStream<E> {
33 receiver: EventStreamReceiver<E>,
34 stop: Arc<AtomicBool>,
35 worker: Option<thread::JoinHandle<()>>,
36}
37enum EventStreamReceiver<E> {
38 Futures(UnboundedReceiver<Result<E>>),
39 Tokio(TokioEventReceiver<E>),
40}
41
42pub fn from_receiver<E>(receiver: EventReceiver<E>) -> EventStream<E>
48where
49 E: Send + 'static,
50{
51 let (sender, async_receiver) = unbounded();
52 let stop = Arc::new(AtomicBool::new(false));
53 let worker_stop = Arc::clone(&stop);
54 let worker = thread::spawn(move || forward_receiver(receiver, sender, worker_stop));
55 EventStream {
56 receiver: EventStreamReceiver::Futures(async_receiver),
57 stop,
58 worker: Some(worker),
59 }
60}
61
62fn forward_receiver<E>(
67 receiver: EventReceiver<E>,
68 sender: futures::channel::mpsc::UnboundedSender<Result<E>>,
69 stop: Arc<AtomicBool>,
70) where
71 E: Send + 'static,
72{
73 while !stop.load(Ordering::Acquire) {
74 match receiver.recv_timeout(Duration::from_millis(50)) {
75 Ok(Some(event)) => {
76 if sender.unbounded_send(Ok(event)).is_err() {
77 break;
78 }
79 }
80 Ok(None) => {}
81 Err(error) => {
82 let _ = sender.unbounded_send(Err(error));
83 break;
84 }
85 }
86 }
87}
88
89pub fn from_tokio_receiver<E>(receiver: TokioEventReceiver<E>) -> EventStream<E> {
91 EventStream {
92 receiver: EventStreamReceiver::Tokio(receiver),
93 stop: Arc::new(AtomicBool::new(false)),
94 worker: None,
95 }
96}
97
98impl<E> Stream for EventStream<E> {
99 type Item = Result<E>;
100
101 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
102 match &mut self.receiver {
103 EventStreamReceiver::Futures(receiver) => Pin::new(receiver).poll_next(cx),
104 EventStreamReceiver::Tokio(receiver) => Pin::new(receiver).poll_recv(cx),
105 }
106 }
107}
108
109impl<E> Drop for EventStream<E> {
110 fn drop(&mut self) {
111 self.stop.store(true, Ordering::Release);
112 if let Some(worker) = self.worker.take() {
113 let _ = worker.join();
114 }
115 }
116}
117
118#[cfg(test)]
119mod tests {
120 use super::*;
121 use futures::StreamExt;
122
123 #[test]
124 fn worker_forwards_events_to_the_stream() {
125 let (sender, receiver) = EventReceiver::bounded();
126 let mut events = from_receiver(receiver);
127 assert!(sender.send(7_u8, 0));
128 assert_eq!(
129 futures::executor::block_on(events.next()).unwrap().unwrap(),
130 7
131 );
132 }
133
134 #[test]
135 fn worker_forwards_receiver_errors() {
136 let (sender, receiver) = EventReceiver::<u8>::bounded();
137 let mut events = from_receiver(receiver);
138 assert!(sender.send_error(Error::InvalidState));
139 assert!(futures::executor::block_on(events.next()).unwrap().is_err());
140 }
141
142 #[test]
143 fn native_tokio_receiver_uses_the_same_stream_surface() {
144 let (sender, receiver) = TokioEventReceiver::bounded();
145 let mut events = from_tokio_receiver(receiver);
146 assert!(sender.send(7_u8, || 0));
147 drop(sender);
148 assert_eq!(
149 futures::executor::block_on(events.next()).unwrap().unwrap(),
150 7
151 );
152 assert!(futures::executor::block_on(events.next()).is_none());
153 }
154
155 #[test]
156 fn dropping_a_sync_stream_joins_its_adapter_worker() {
157 let (_sender, receiver) = EventReceiver::<u8>::bounded();
158 drop(from_receiver(receiver));
159 }
160
161 #[test]
162 fn worker_stops_when_its_async_consumer_is_already_gone() {
163 let (input, receiver) = EventReceiver::bounded();
164 let (output, async_receiver) = unbounded::<Result<u8>>();
165 drop(async_receiver);
166 assert!(input.send(7, 0));
167 forward_receiver(receiver, output, Arc::new(AtomicBool::new(false)));
168 }
169
170 #[test]
171 fn worker_rechecks_an_empty_receiver_until_shutdown_is_requested() {
172 let (_sender, receiver) = EventReceiver::<u8>::bounded();
173 let (output, _async_receiver) = unbounded::<Result<u8>>();
174 let stop = Arc::new(AtomicBool::new(false));
175 let stop_later = Arc::clone(&stop);
176 std::thread::spawn(move || {
177 std::thread::sleep(Duration::from_millis(75));
178 stop_later.store(true, Ordering::Release);
179 });
180 forward_receiver(receiver, output, stop);
181 }
182
183 #[test]
184 fn adapter_handles_immediate_shutdown_and_native_stream_without_worker() {
185 let (_sender, receiver) = EventReceiver::<u8>::bounded();
186 let (output, _async_receiver) = unbounded::<Result<u8>>();
187 forward_receiver(receiver, output, Arc::new(AtomicBool::new(true)));
188
189 let (_sender, receiver) = TokioEventReceiver::<u8>::bounded();
190 drop(from_tokio_receiver(receiver));
191 }
192
193 #[test]
194 fn native_stream_drop_does_not_attempt_to_join_a_worker() {
195 let (_sender, receiver) = TokioEventReceiver::<u8>::bounded();
196 drop(from_tokio_receiver(receiver));
197 }
198
199 #[test]
200 fn dropping_a_stream_ignores_a_worker_join_error() {
201 let (_sender, receiver) = TokioEventReceiver::<u8>::bounded();
202 let worker = std::thread::spawn(|| panic!("worker terminated unexpectedly"));
203 let stream = EventStream {
204 receiver: EventStreamReceiver::Tokio(receiver),
205 stop: Arc::new(AtomicBool::new(false)),
206 worker: Some(worker),
207 };
208 drop(stream);
209 }
210}