Skip to main content

dnet_rpc/consumer/
stream.rs

1//! Return value for streaming requests.
2
3use std::{
4    pin::Pin,
5    sync::{Arc, Mutex},
6    task::{Context, Poll},
7};
8
9use futures::{
10    channel::{
11        mpsc::{unbounded, UnboundedReceiver},
12        oneshot,
13    },
14    future::FusedFuture,
15    ready,
16    stream::FusedStream,
17    Future, FutureExt, StreamExt,
18};
19use pin_project::{pin_project, pinned_drop};
20
21use crate::parts::consumer::{RequestSender, ResultSender};
22
23use super::Aborter;
24
25/// Future returned by consumer streaming requests.
26#[derive(Debug)]
27#[pin_project(PinnedDrop)]
28pub struct StreamRequest<Request, T> {
29    sender: RequestSender<Request, T>,
30    id: u64,
31    request: Option<Request>,
32    result_receiver: Option<oneshot::Receiver<super::Result<()>>>,
33    values_receiver: Option<UnboundedReceiver<T>>,
34    aborter: Option<Aborter<Request, T>>,
35    abort_receiver: Option<oneshot::Receiver<()>>,
36}
37
38impl<Request, T> StreamRequest<Request, T> {
39    /// Create new stream request.
40    ///
41    /// **NOTE**: This is used internally by the generated consumers.<br>
42    /// You should never have to create it manually yourself.
43    pub fn new(sender: RequestSender<Request, T>, id: u64, request: Request) -> Self {
44        StreamRequest {
45            sender,
46            id,
47            request: Some(request),
48            result_receiver: None,
49            values_receiver: None,
50            aborter: None,
51            abort_receiver: None,
52        }
53    }
54
55    /// Request id.
56    pub fn id(&self) -> u64 {
57        self.id
58    }
59
60    /// Aborter for this stream request.
61    ///
62    /// **NOTE**: This aborter has ability to abort resulting stream as well.
63    pub fn aborter(&mut self) -> Aborter<Request, T> {
64        let aborter = self.aborter.get_or_insert_with(|| {
65            let (abort_sender, abort_receiver) = oneshot::channel();
66            self.abort_receiver = Some(abort_receiver);
67            Aborter {
68                id: self.id,
69                sender: self.sender.clone(),
70                abort_sender: Arc::new(Mutex::new(Some(abort_sender))),
71            }
72        });
73        aborter.clone()
74    }
75}
76
77impl<Request, T> Future for StreamRequest<Request, T> {
78    type Output = super::Result<Stream<Request, T>>;
79
80    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
81        let mut me = self.project();
82        if let Some(abort_receiver) = &mut me.abort_receiver {
83            if let Poll::Ready(result) = abort_receiver.poll_unpin(cx) {
84                if result.is_ok() {
85                    me.request.take();
86                    me.result_receiver.take();
87                    me.values_receiver.take();
88                    return Poll::Ready(Err(super::super::Error::Aborted));
89                }
90            }
91        }
92
93        if let Some(request) = me.request.take() {
94            let (result_sender, mut result_receiver) = oneshot::channel();
95            let (values_sender, values_receiver) = unbounded();
96            let sender = ResultSender::Stream {
97                result_sender,
98                values_sender,
99            };
100            me.sender.send(*me.id, request, sender).map_err(|_| {
101                me.result_receiver.take();
102                super::super::Error::Shutdown
103            })?;
104            match result_receiver.poll_unpin(cx) {
105                Poll::Ready(result) => {
106                    me.result_receiver.take();
107                    let result = result
108                        .map(|_| Stream {
109                            id: *me.id,
110                            sender: me.sender.clone(),
111                            receiver: Some(values_receiver),
112                            aborter: me.aborter.take(),
113                            abort_receiver: me.abort_receiver.take(),
114                        })
115                        .map_err(|_| {
116                            me.values_receiver.take();
117                            super::super::Error::Dropped
118                        });
119                    Poll::Ready(result)
120                }
121                Poll::Pending => {
122                    *me.result_receiver = Some(result_receiver);
123                    *me.values_receiver = Some(values_receiver);
124                    Poll::Pending
125                }
126            }
127        } else if let Some(receiver) = &mut me.result_receiver {
128            let result = ready!(receiver.poll_unpin(cx));
129            me.result_receiver.take();
130            let result = result
131                .map(|_| Stream {
132                    id: *me.id,
133                    sender: me.sender.clone(),
134                    receiver: me.values_receiver.take(),
135                    aborter: me.aborter.take(),
136                    abort_receiver: me.abort_receiver.take(),
137                })
138                .map_err(|_| {
139                    me.values_receiver.take();
140                    super::super::Error::Dropped
141                });
142            Poll::Ready(result)
143        } else {
144            Poll::Pending
145        }
146    }
147}
148
149impl<Request, T> FusedFuture for StreamRequest<Request, T> {
150    fn is_terminated(&self) -> bool {
151        self.request.is_none() && self.result_receiver.is_none()
152    }
153}
154
155#[pinned_drop]
156impl<Request, T> PinnedDrop for StreamRequest<Request, T> {
157    fn drop(self: Pin<&mut Self>) {
158        if !self.is_terminated() {
159            self.sender.abort(self.id);
160        }
161    }
162}
163
164/// Stream returned by the [StreamRequest].
165#[derive(Debug)]
166pub struct Stream<Request, T> {
167    id: u64,
168    sender: RequestSender<Request, T>,
169    receiver: Option<UnboundedReceiver<T>>,
170    aborter: Option<Aborter<Request, T>>,
171    abort_receiver: Option<oneshot::Receiver<()>>,
172}
173
174impl<Request, T> Stream<Request, T> {
175    /// Aborter for this stream.
176    pub fn aborter(&mut self) -> Aborter<Request, T> {
177        let aborter = self.aborter.get_or_insert_with(|| {
178            let (abort_sender, abort_receiver) = oneshot::channel();
179            self.abort_receiver = Some(abort_receiver);
180            Aborter {
181                id: self.id,
182                sender: self.sender.clone(),
183                abort_sender: Arc::new(Mutex::new(Some(abort_sender))),
184            }
185        });
186        aborter.clone()
187    }
188}
189
190impl<Request, T> futures::Stream for Stream<Request, T> {
191    type Item = T;
192
193    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
194        if let Some(receiver) = &mut self.receiver {
195            match receiver.poll_next_unpin(cx) {
196                Poll::Ready(result) => {
197                    if let Some(result) = result {
198                        Poll::Ready(Some(result))
199                    } else {
200                        self.receiver.take();
201                        Poll::Ready(None)
202                    }
203                }
204                Poll::Pending => {
205                    if let Some(abort_receiver) = &mut self.abort_receiver {
206                        if let Poll::Ready(result) = abort_receiver.poll_unpin(cx) {
207                            if result.is_ok() {
208                                self.receiver.take();
209                                return Poll::Ready(None);
210                            }
211                        }
212                    }
213                    Poll::Pending
214                }
215            }
216        } else {
217            Poll::Ready(None)
218        }
219    }
220}
221
222impl<Request, T> FusedStream for Stream<Request, T> {
223    fn is_terminated(&self) -> bool {
224        self.receiver.is_none()
225    }
226}
227
228impl<Request, T> Drop for Stream<Request, T> {
229    fn drop(&mut self) {
230        if !self.is_terminated() {
231            self.sender.abort(self.id);
232        }
233    }
234}