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 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 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 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 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
176pub 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 <[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
323pub trait SubmitMultiFactory {
331 type Op: OpCode + TakeBuffer<Buffer = Self::Buffer> + 'static;
333
334 type Buffer: HandleBufferRef + 'static;
339
340 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
358pub struct SubmitMultiStream<F: SubmitMultiFactory> {
361 factory: F,
362 op: Option<SubmitMultiManaged<F::Op, F::Buffer>>,
363}
364
365impl<F: SubmitMultiFactory> SubmitMultiStream<F> {
366 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}