1use 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#[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 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 pub fn id(&self) -> u64 {
57 self.id
58 }
59
60 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#[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 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}