Skip to main content

compio_runtime/future/
stream.rs

1use std::{
2    cell::RefCell,
3    io,
4    marker::PhantomData,
5    pin::Pin,
6    rc::Rc,
7    task::{Context, Poll},
8};
9
10use compio_buf::{BufResult, SetLenExt};
11use compio_driver::{
12    BufferPool, BufferRef, Extra, Key, OpCode, Proactor, PushEntry, TakeBuffer,
13    op::{RecvFromMultiResult, RecvMsgMultiResult},
14};
15use futures_util::{Stream, StreamExt, stream::FusedStream};
16
17use crate::{
18    CancelToken, ContextExt,
19    future::{poll_multishot, poll_task_with_extra, submit_raw},
20};
21
22pin_project_lite::pin_project! {
23    /// Returned [`Stream`] for [`Runtime::submit_multi`].
24    ///
25    /// When this is dropped and the operation hasn't finished yet, it will try to
26    /// cancel the operation.
27    pub struct SubmitMulti<T: OpCode> {
28        driver: Rc<RefCell<Proactor>>,
29        state: Option<State<T>>,
30    }
31
32    impl<T: OpCode> PinnedDrop for SubmitMulti<T> {
33        fn drop(this: Pin<&mut Self>) {
34            let this = this.project();
35            if let Some(State::Submitted { key }) = this.state.take() {
36                this.driver.borrow_mut().cancel(key);
37            }
38        }
39    }
40}
41
42enum State<T: OpCode> {
43    Idle { op: T },
44    Submitted { key: Key<T> },
45    Finished { op: T },
46}
47
48impl<T: OpCode> State<T> {
49    fn submitted(key: Key<T>) -> Self {
50        State::Submitted { key }
51    }
52}
53
54impl<T: OpCode> SubmitMulti<T> {
55    pub(crate) fn new(driver: Rc<RefCell<Proactor>>, op: T) -> Self {
56        SubmitMulti {
57            driver,
58            state: Some(State::Idle { op }),
59        }
60    }
61
62    /// Try to take the inner op from the stream.
63    ///
64    /// Returns `Ok(T)` if the stream:
65    ///
66    /// - has not been polled yet, or
67    /// - is finished and the op is returned by the driver
68    ///
69    /// Returns `Err(Self)` if it's still running.
70    pub fn try_take(mut self) -> Result<T, Self> {
71        match self.state.take() {
72            Some(State::Finished { op }) | Some(State::Idle { op }) => Ok(op),
73            state => {
74                debug_assert!(state.is_some());
75                self.state = state;
76                Err(self)
77            }
78        }
79    }
80}
81
82impl<T: OpCode + 'static> Stream for SubmitMulti<T> {
83    type Item = BufResult<usize, Extra>;
84
85    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
86        let this = self.project();
87
88        loop {
89            match this.state.take().expect("State error, this is a bug") {
90                State::Idle { op } => {
91                    let extra = cx.as_extra(|| this.driver.borrow().default_extra());
92                    let entry = submit_raw(&mut this.driver.borrow_mut(), op, extra);
93                    match entry {
94                        PushEntry::Pending(key) => {
95                            if let Some(cancel) = cx.get_cancel() {
96                                cancel.register(&key);
97                            }
98
99                            *this.state = Some(State::submitted(key))
100                        }
101                        PushEntry::Ready(BufResult(res, op)) => {
102                            *this.state = Some(State::Finished { op });
103                            let extra = this.driver.borrow().default_extra();
104
105                            return Poll::Ready(Some(BufResult(res, extra)));
106                        }
107                    }
108                }
109
110                State::Submitted { key, .. } => {
111                    if let Some(res) =
112                        poll_multishot(&mut this.driver.borrow_mut(), cx.get_waker(), &key)
113                    {
114                        *this.state = Some(State::submitted(key));
115
116                        return Poll::Ready(Some(res));
117                    };
118
119                    let entry =
120                        poll_task_with_extra(&mut this.driver.borrow_mut(), cx.get_waker(), key);
121                    match entry {
122                        PushEntry::Pending(key) => {
123                            *this.state = Some(State::submitted(key));
124
125                            return Poll::Pending;
126                        }
127                        PushEntry::Ready((BufResult(res, op), extra)) => {
128                            *this.state = Some(State::Finished { op });
129
130                            return Poll::Ready(Some(BufResult(res, extra)));
131                        }
132                    }
133                }
134
135                State::Finished { op } => {
136                    *this.state = Some(State::Finished { op });
137
138                    return Poll::Ready(None);
139                }
140            }
141        }
142    }
143}
144
145impl<T: OpCode + 'static> FusedStream for SubmitMulti<T> {
146    fn is_terminated(&self) -> bool {
147        matches!(self.state, None | Some(State::Finished { .. }))
148    }
149}
150
151impl<T: OpCode + TakeBuffer + 'static> SubmitMulti<T>
152where
153    <T as TakeBuffer>::Buffer: HandleBufferRef<Param = ()>,
154{
155    /// Convert this stream into one that iterates the buffers from the results.
156    pub fn into_managed(self, buffer_pool: BufferPool) -> SubmitMultiManaged<T, T::Buffer> {
157        SubmitMultiManaged::new(self, buffer_pool, ())
158    }
159}
160
161impl<T: OpCode + TakeBuffer + 'static> SubmitMulti<T>
162where
163    <T as TakeBuffer>::Buffer: HandleBufferRef,
164{
165    /// Convert this stream into one that iterates the buffers from the results,
166    /// with a param to construct the result item.
167    pub fn into_managed_with(
168        self,
169        buffer_pool: BufferPool,
170        param: <<T as TakeBuffer>::Buffer as HandleBufferRef>::Param,
171    ) -> SubmitMultiManaged<T, T::Buffer> {
172        SubmitMultiManaged::new(self, buffer_pool, param)
173    }
174}
175
176/// A wrapper around [`SubmitMulti`] that iterates the buffers from the results.
177pub struct SubmitMultiManaged<T: OpCode, B = BufferRef>
178where
179    B: HandleBufferRef + 'static,
180{
181    inner: Option<SubmitMulti<T>>,
182    buffer_pool: BufferPool,
183    param: <B as HandleBufferRef>::Param,
184    _p: PhantomData<&'static B>,
185}
186
187impl<T: OpCode, B: HandleBufferRef + 'static> SubmitMultiManaged<T, B> {
188    fn new(
189        stream: SubmitMulti<T>,
190        buffer_pool: BufferPool,
191        param: <B as HandleBufferRef>::Param,
192    ) -> Self {
193        Self {
194            inner: Some(stream),
195            buffer_pool,
196            param,
197            _p: PhantomData,
198        }
199    }
200}
201
202impl<T: OpCode + TakeBuffer<Buffer = B> + 'static, B: HandleBufferRef> Stream
203    for SubmitMultiManaged<T, B>
204{
205    type Item = io::Result<Option<B>>;
206
207    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
208        if let Some(inner) = self.inner.as_mut() {
209            let buffer = match std::task::ready!(inner.poll_next_unpin(cx)) {
210                Some(BufResult(res, extra)) => {
211                    if inner.is_terminated() {
212                        let mut b = self
213                            .inner
214                            .take()
215                            .and_then(|s| s.try_take().ok())
216                            .and_then(|op| op.take_buffer());
217                        let res = res?;
218                        if let Some(ref mut b) = b {
219                            unsafe { b.advance_to(res) }
220                        }
221                        b
222                    } else {
223                        let b = self.buffer_pool.take(extra.buffer_id()?)?;
224                        let res = res?;
225                        if let Some(mut b) = b {
226                            unsafe {
227                                SetLenExt::advance_to(&mut b, res);
228                                Some(B::from_buffer_ref(b, self.param))
229                            }
230                        } else {
231                            None
232                        }
233                    }
234                }
235                None => self
236                    .inner
237                    .take()
238                    .and_then(|s| s.try_take().ok())
239                    .and_then(|op| op.take_buffer()),
240            };
241            Poll::Ready(Some(Ok(buffer)))
242        } else {
243            Poll::Ready(None)
244        }
245    }
246}
247
248impl<T: OpCode + TakeBuffer<Buffer = B> + 'static, B: HandleBufferRef> FusedStream
249    for SubmitMultiManaged<T, B>
250{
251    fn is_terminated(&self) -> bool {
252        self.inner.as_ref().is_none_or(|s| s.is_terminated())
253    }
254}
255
256mod private {
257    use super::*;
258
259    pub trait Sealed {}
260
261    impl Sealed for BufferRef {}
262    impl Sealed for RecvFromMultiResult {}
263    impl Sealed for RecvMsgMultiResult {}
264}
265
266#[doc(hidden)]
267pub trait HandleBufferRef: private::Sealed {
268    type Param: Copy + Unpin;
269
270    unsafe fn from_buffer_ref(buffer: BufferRef, param: Self::Param) -> Self;
271
272    unsafe fn advance_to(&mut self, len: usize);
273
274    fn is_empty(&self) -> bool;
275}
276
277impl HandleBufferRef for BufferRef {
278    type Param = ();
279
280    unsafe fn from_buffer_ref(buffer: BufferRef, _: Self::Param) -> Self {
281        buffer
282    }
283
284    unsafe fn advance_to(&mut self, len: usize) {
285        unsafe { SetLenExt::advance_to(self, len) }
286    }
287
288    fn is_empty(&self) -> bool {
289        // A fallback buffer pool takes the buffer before the operation, so it
290        // can return an empty buffer when EOF is reached.
291        <[u8]>::is_empty(self)
292    }
293}
294
295impl HandleBufferRef for RecvFromMultiResult {
296    type Param = ();
297
298    unsafe fn from_buffer_ref(buffer: BufferRef, _: Self::Param) -> Self {
299        unsafe { RecvFromMultiResult::new(buffer) }
300    }
301
302    unsafe fn advance_to(&mut self, _: usize) {}
303
304    fn is_empty(&self) -> bool {
305        false
306    }
307}
308
309impl HandleBufferRef for RecvMsgMultiResult {
310    type Param = usize;
311
312    unsafe fn from_buffer_ref(buffer: BufferRef, clen: usize) -> Self {
313        unsafe { RecvMsgMultiResult::new(buffer, clen) }
314    }
315
316    unsafe fn advance_to(&mut self, _: usize) {}
317
318    fn is_empty(&self) -> bool {
319        false
320    }
321}
322
323/// Creates managed multishot submissions for [`SubmitMultiStream`].
324///
325/// [`SubmitMultiStream`] calls [`create`](Self::create) again when the previous
326/// submission ends before EOF.
327///
328/// By default, this trait is implemented for closures
329/// `FnMut() -> std::io::Result<SubmitMultiManaged<T, B>>`.
330pub trait SubmitMultiFactory {
331    /// The [`OpCode`] type.
332    type Op: OpCode + TakeBuffer<Buffer = Self::Buffer> + 'static;
333
334    /// The buffer returned by the OpCode.
335    ///
336    /// This can be [`BufferRef`], [`RecvFromMultiResult`] or
337    /// [`RecvMsgMultiResult`].
338    type Buffer: HandleBufferRef + 'static;
339
340    /// Creates a new managed multishot submission.
341    fn create(&mut self) -> io::Result<SubmitMultiManaged<Self::Op, Self::Buffer>>;
342}
343
344impl<F, T, B> SubmitMultiFactory for F
345where
346    F: FnMut() -> io::Result<SubmitMultiManaged<T, B>>,
347    T: OpCode + TakeBuffer<Buffer = B> + 'static,
348    B: HandleBufferRef + 'static,
349{
350    type Buffer = B;
351    type Op = T;
352
353    fn create(&mut self) -> io::Result<SubmitMultiManaged<Self::Op, Self::Buffer>> {
354        (self)()
355    }
356}
357
358/// A wrapper around [`SubmitMultiManaged`] that submits operations
359/// automatically until the stream is finished.
360pub struct SubmitMultiStream<F: SubmitMultiFactory> {
361    factory: F,
362    op: Option<SubmitMultiManaged<F::Op, F::Buffer>>,
363}
364
365impl<F: SubmitMultiFactory> SubmitMultiStream<F> {
366    /// Creates a new [`SubmitMultiStream`] with `factory`.
367    pub fn new(factory: F) -> Self {
368        Self { factory, op: None }
369    }
370}
371
372impl<F> Stream for SubmitMultiStream<F>
373where
374    F: SubmitMultiFactory + Unpin,
375{
376    type Item = io::Result<F::Buffer>;
377
378    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
379        loop {
380            match &mut self.op {
381                Some(op) => match std::task::ready!(Pin::new(op).poll_next(cx)) {
382                    Some(Ok(Some(buffer))) => {
383                        if buffer.is_empty() {
384                            break Poll::Ready(None);
385                        } else {
386                            break Poll::Ready(Some(Ok(buffer)));
387                        }
388                    }
389                    Some(Ok(None)) => break Poll::Ready(None),
390                    Some(Err(e)) => break Poll::Ready(Some(Err(e))),
391                    None => self.op = None,
392                },
393                None if cx.get_cancel().is_some_and(CancelToken::is_cancelled) => {
394                    break Poll::Ready(None);
395                }
396                None => match self.factory.create() {
397                    Ok(op) => self.op = Some(op),
398                    Err(e) => break Poll::Ready(Some(Err(e))),
399                },
400            }
401        }
402    }
403}