Skip to main content

ntex_net/tokio/
io.rs

1use std::task::{Context, Poll, ready};
2use std::{any, cmp, future::poll_fn, io, mem, pin::Pin, ptr, rc::Rc};
3
4use ntex_bytes::{BufMut, BytePage};
5use ntex_io::{Filter, Handle, Io, IoBoxed, IoContext, IoStream, IoTaskStatus, Readiness, types};
6use tok_io::io::{AsyncRead, AsyncWrite, ReadBuf};
7use tok_io::net::TcpStream;
8
9impl IoStream for super::TcpStream {
10    fn start(self, ctx: IoContext) -> Box<dyn Handle> {
11        let io = Rc::new(self.0);
12        tok_io::task::spawn_local(run_rd(io.clone(), ctx.clone()));
13        tok_io::task::spawn_local(run_wrt(io.clone(), ctx));
14        Box::new(HandleWrapper(io))
15    }
16}
17
18#[cfg(unix)]
19impl IoStream for super::UnixStream {
20    fn start(self, ctx: IoContext) -> Box<dyn Handle> {
21        let io = Rc::new(self.0);
22        tok_io::task::spawn_local(run_rd(io.clone(), ctx.clone()));
23        tok_io::task::spawn_local(run_wrt(io.clone(), ctx));
24        Box::new(HandleWrapperUnix(io))
25    }
26}
27
28trait Stream: AsyncRead + AsyncWrite + Unpin {
29    fn poll_read_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>>;
30
31    fn poll_write_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>>;
32
33    /// Closes both directions gracefully, after draining the receive queue.
34    fn terminate(&self) -> io::Result<()>;
35
36    /// Arranges for the socket to be reset instead of closed gracefully.
37    fn abort(&self);
38
39    fn try_read(&self, buf: &mut [u8]) -> io::Result<usize>;
40
41    fn try_write(&self, buf: &[u8]) -> io::Result<usize>;
42
43    fn try_write_vectored(&self, buf: &[io::IoSlice<'_>]) -> io::Result<usize>;
44}
45
46impl Stream for TcpStream {
47    fn poll_read_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
48        TcpStream::poll_read_ready(self, cx)
49    }
50
51    fn poll_write_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
52        TcpStream::poll_write_ready(self, cx)
53    }
54
55    fn terminate(&self) -> io::Result<()> {
56        let sock = socket2::SockRef::from(self);
57        crate::helpers::drain_socket(&sock);
58        crate::helpers::shutdown_result(sock.shutdown(std::net::Shutdown::Both))
59    }
60
61    fn abort(&self) {
62        crate::helpers::abort_socket(&socket2::SockRef::from(self));
63    }
64
65    fn try_read(&self, buf: &mut [u8]) -> io::Result<usize> {
66        TcpStream::try_read(self, buf)
67    }
68
69    fn try_write(&self, buf: &[u8]) -> io::Result<usize> {
70        TcpStream::try_write(self, buf)
71    }
72
73    fn try_write_vectored(&self, buf: &[io::IoSlice<'_>]) -> io::Result<usize> {
74        TcpStream::try_write_vectored(self, buf)
75    }
76}
77
78#[cfg(unix)]
79impl Stream for tok_io::net::UnixStream {
80    fn poll_read_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
81        tok_io::net::UnixStream::poll_read_ready(self, cx)
82    }
83
84    fn poll_write_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
85        tok_io::net::UnixStream::poll_write_ready(self, cx)
86    }
87
88    fn terminate(&self) -> io::Result<()> {
89        let sock = socket2::SockRef::from(self);
90        crate::helpers::drain_socket(&sock);
91        crate::helpers::shutdown_result(sock.shutdown(std::net::Shutdown::Both))
92    }
93
94    fn abort(&self) {
95        crate::helpers::abort_socket(&socket2::SockRef::from(self));
96    }
97
98    fn try_read(&self, buf: &mut [u8]) -> io::Result<usize> {
99        tok_io::net::UnixStream::try_read(self, buf)
100    }
101
102    fn try_write(&self, buf: &[u8]) -> io::Result<usize> {
103        tok_io::net::UnixStream::try_write(self, buf)
104    }
105
106    fn try_write_vectored(&self, buf: &[io::IoSlice<'_>]) -> io::Result<usize> {
107        tok_io::net::UnixStream::try_write_vectored(self, buf)
108    }
109}
110
111struct HandleWrapper(Rc<TcpStream>);
112
113impl Handle for HandleWrapper {
114    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
115        if id == any::TypeId::of::<types::PeerAddr>() {
116            let result = self.0.peer_addr();
117            if let Ok(addr) = result {
118                return Some(Box::new(types::PeerAddr(addr)));
119            }
120        }
121        None
122    }
123
124    fn write(&self, ctx: &IoContext) {
125        let _ = write(self.0.as_ref(), ctx, true);
126    }
127}
128
129#[cfg(unix)]
130struct HandleWrapperUnix(Rc<tok_io::net::UnixStream>);
131
132#[cfg(unix)]
133impl Handle for HandleWrapperUnix {
134    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
135        None
136    }
137
138    fn write(&self, ctx: &IoContext) {
139        let _ = write(self.0.as_ref(), ctx, true);
140    }
141}
142
143async fn run_rd<T>(io: Rc<T>, ctx: IoContext)
144where
145    T: Stream + Unpin,
146{
147    let st = poll_fn(|cx| {
148        'outer: loop {
149            let ctx_state = ctx.poll_read_ready(cx);
150            #[cfg(feature = "trace")]
151            log::trace!(
152                "{}: Read task, ctx:{ctx_state:?} flags:{:?}",
153                ctx.tag(),
154                ctx.flags()
155            );
156            return match ready!(ctx_state) {
157                Readiness::Ready => {
158                    let io_state = io.poll_read_ready(cx);
159                    #[cfg(feature = "trace")]
160                    log::trace!("{}: Io read ready: {io_state:?}", ctx.tag());
161
162                    match ready!(io_state) {
163                        Ok(()) => 'inner: loop {
164                            return match read(io.as_ref(), &ctx) {
165                                Poll::Ready(IoTaskStatus::Io) => continue 'inner,
166                                Poll::Ready(IoTaskStatus::Pause) => Poll::Pending,
167                                Poll::Ready(IoTaskStatus::Stop) => Poll::Ready(()),
168                                Poll::Pending => continue 'outer,
169                            };
170                        },
171                        Err(err) => {
172                            ctx.stop(Some(err));
173                            Poll::Ready(())
174                        }
175                    }
176                }
177                Readiness::Close | Readiness::Terminate => Poll::Ready(()),
178            };
179        }
180    })
181    .await;
182}
183
184#[derive(Copy, Clone, PartialEq, Eq, Debug)]
185enum WrtStatus {
186    More,
187    Pending,
188    Terminate,
189}
190
191async fn run_wrt<T>(io: Rc<T>, ctx: IoContext)
192where
193    T: Stream,
194{
195    let terminate = poll_fn(|cx| {
196        let ctx_state = ctx.poll_write_ready(cx);
197        #[cfg(feature = "trace")]
198        log::trace!(
199            "{}: Write task, ctx {ctx_state:?} flags:{:?}",
200            ctx.tag(),
201            ctx.flags()
202        );
203
204        match ready!(ctx_state) {
205            Readiness::Ready => loop {
206                let io_state = io.poll_write_ready(cx);
207                #[cfg(feature = "trace")]
208                log::trace!("{}: Io write ready {io_state:?}", ctx.tag());
209
210                return match ready!(io_state) {
211                    Ok(()) => match write(io.as_ref(), &ctx, false) {
212                        WrtStatus::More => continue,
213                        WrtStatus::Pending => Poll::Pending,
214                        // the connection has been aborted already
215                        WrtStatus::Terminate => Poll::Ready(true),
216                    },
217                    Err(err) => {
218                        ctx.update_write_status(Err(err));
219                        Poll::Ready(true)
220                    }
221                };
222            },
223            Readiness::Close => Poll::Ready(false),
224            Readiness::Terminate => Poll::Ready(true),
225        }
226    })
227    .await;
228
229    log::trace!("{}: Shuting down io", ctx.tag());
230
231    // A force-closed connection is aborted instead of closed gracefully, so
232    // that a truncated stream is not terminated by a clean `FIN`.
233    let result = if terminate {
234        io.abort();
235        Ok(())
236    } else {
237        io.terminate()
238    };
239
240    log::trace!("{}: Shutdown complete {result:?}", ctx.tag());
241    ctx.stopped(result.err());
242}
243
244const MAX_WRITE_SIZE: usize = 64 * 1024;
245const MAX_WRITE_ITEMS: usize = 16;
246
247fn write<T>(io: &T, ctx: &IoContext, direct: bool) -> WrtStatus
248where
249    T: Stream,
250{
251    let result = ctx.with_write_dst(|dst| {
252        let mut pages: [Option<BytePage>; MAX_WRITE_ITEMS] = [
253            None, None, None, None, None, None, None, None, None, None, None, None, None, None,
254            None, None,
255        ];
256        let mut bufs: [mem::MaybeUninit<io::IoSlice<'_>>; MAX_WRITE_ITEMS] =
257            [mem::MaybeUninit::uninit(); MAX_WRITE_ITEMS];
258
259        let mut num = 0;
260        let mut size = 0;
261
262        #[cfg(feature = "trace")]
263        log::trace!("{}: Try write buf({})", ctx.tag(), dst.len());
264
265        while let Some(page) = dst.take() {
266            size += page.len();
267
268            // The slice is taken from the stored page, an inline page keeps
269            // its data in the `BytePage` itself and moving it moves the data.
270            // SAFETY: Page is stored in `pages` for lifetime of `bufs` and is
271            // not moved until `bufs` is dropped
272            let page = pages[num].insert(page);
273            bufs[num] = mem::MaybeUninit::new(io::IoSlice::new(unsafe {
274                mem::transmute::<&[u8], &[u8]>(page.as_ref())
275            }));
276
277            num += 1;
278            if num == MAX_WRITE_ITEMS || size >= MAX_WRITE_SIZE {
279                break;
280            }
281        }
282
283        if num > 0 {
284            // SAFETY: initialize in previous block
285            let bufs = unsafe { &*(&raw const bufs[..num] as *const [std::io::IoSlice<'_>]) };
286
287            // An error must not return early: the pages taken above still
288            // have to go back to the write buffer.
289            let result = write_io(ctx, io, bufs);
290
291            // remove written bytes
292            if let Poll::Ready(Ok(mut written)) = result {
293                for page in pages[..num].iter_mut().flatten() {
294                    let len = cmp::min(page.len(), written);
295                    page.advance_to(len);
296                    written -= len;
297                    if written == 0 {
298                        break;
299                    }
300                }
301            }
302            // return unwritten data back to the buffer
303            for p in pages[..num].iter_mut().rev() {
304                if let Some(page) = p.take() {
305                    dst.prepend(page);
306                }
307            }
308
309            #[cfg(feature = "trace")]
310            log::trace!(
311                "{}: Io write (direct:{direct}) result:{result:?} buf:{} flags:{:?}",
312                ctx.tag(),
313                dst.len(),
314                ctx.flags()
315            );
316
317            match result? {
318                Poll::Ready(val) => {
319                    if val == 0 {
320                        ctx.stop(None);
321                    }
322                    Ok(val)
323                }
324                Poll::Pending => Ok(0),
325            }
326        } else {
327            Ok(0)
328        }
329    });
330
331    let st = ctx.update_write_status(result);
332    let result = match st {
333        IoTaskStatus::Stop => WrtStatus::Terminate,
334        IoTaskStatus::Pause => WrtStatus::Pending,
335        IoTaskStatus::Io => WrtStatus::More,
336    };
337
338    #[cfg(feature = "trace")]
339    log::trace!(
340        "{}: Write status \"{st:?}({result:?})\" flags:{:?}",
341        ctx.tag(),
342        ctx.flags()
343    );
344    result
345}
346
347/// Flush write buffer to underlying I/O stream.
348fn write_io<T: Stream>(
349    ctx: &IoContext,
350    io: &T,
351    bufs: &[io::IoSlice<'_>],
352) -> Poll<io::Result<usize>> {
353    let result = if bufs.len() == 1 {
354        io.try_write(&bufs[0])
355    } else {
356        io.try_write_vectored(bufs)
357    };
358    match result {
359        Ok(0) => Poll::Ready(Err(io::Error::new(
360            io::ErrorKind::WriteZero,
361            "failed to write frame to transport",
362        ))),
363        Ok(n) => {
364            #[cfg(feature = "trace")]
365            log::trace!("{}: Flushed {n} bytes from {} pages", ctx.tag(), bufs.len());
366            Poll::Ready(Ok(n))
367        }
368        Err(e) if e.kind() == io::ErrorKind::WouldBlock => Poll::Pending,
369        Err(e) => Poll::Ready(Err(e)),
370    }
371}
372
373fn read<T: Stream + Unpin>(io: &T, ctx: &IoContext) -> Poll<IoTaskStatus> {
374    let mut pending = false;
375
376    let result = ctx.with_read_buf(|buf| {
377        #[cfg(feature = "trace")]
378        log::trace!(
379            "{}: Read attempt, buf len({}) cap({})",
380            ctx.tag(),
381            buf.len(),
382            buf.remaining_mut()
383        );
384
385        // read data from socket
386        let io_res = io.try_read(unsafe { &mut *(ptr::from_mut(buf.chunk_mut()) as *mut [u8]) });
387
388        match io_res {
389            Ok(0) => Poll::Ready(Ok(0)),
390            Ok(n) => {
391                // Safety: This is guaranteed to be the number of initialized
392                // bytes due to the invariants provided by `try_read()`.
393                unsafe { buf.advance_mut(n) };
394                Poll::Ready(Ok(n))
395            }
396            Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
397                pending = true;
398                Poll::Pending
399            }
400            Err(e) => Poll::Ready(Err(e)),
401        }
402    });
403
404    #[cfg(feature = "trace")]
405    log::trace!(
406        "{}: Read status \"{result:?}\" pending({pending})",
407        ctx.tag()
408    );
409
410    if result == IoTaskStatus::Io && pending {
411        Poll::Pending
412    } else {
413        Poll::Ready(result)
414    }
415}
416
417#[derive(Debug)]
418pub struct TokioIoBoxed(IoBoxed);
419
420impl std::ops::Deref for TokioIoBoxed {
421    type Target = IoBoxed;
422
423    #[inline]
424    fn deref(&self) -> &Self::Target {
425        &self.0
426    }
427}
428
429impl From<IoBoxed> for TokioIoBoxed {
430    fn from(io: IoBoxed) -> TokioIoBoxed {
431        TokioIoBoxed(io)
432    }
433}
434
435impl<F: Filter> From<Io<F>> for TokioIoBoxed {
436    fn from(io: Io<F>) -> TokioIoBoxed {
437        TokioIoBoxed(IoBoxed::from(io))
438    }
439}
440
441impl AsyncRead for TokioIoBoxed {
442    fn poll_read(
443        self: Pin<&mut Self>,
444        cx: &mut Context<'_>,
445        buf: &mut ReadBuf<'_>,
446    ) -> Poll<io::Result<()>> {
447        let len = self.0.with_read_dst(|src| {
448            let len = cmp::min(src.len(), buf.remaining());
449            buf.put_slice(&src.split_to(len));
450            len
451        });
452
453        if len == 0 {
454            match ready!(self.0.poll_read_more(cx)) {
455                Ok(Some(())) => Poll::Pending,
456                Err(e) => Poll::Ready(Err(e)),
457                Ok(None) => Poll::Ready(Ok(())),
458            }
459        } else {
460            Poll::Ready(Ok(()))
461        }
462    }
463}
464
465impl AsyncWrite for TokioIoBoxed {
466    fn poll_write(
467        self: Pin<&mut Self>,
468        _: &mut Context<'_>,
469        buf: &[u8],
470    ) -> Poll<io::Result<usize>> {
471        self.0.encode_slice(buf)?;
472        Poll::Ready(Ok(buf.len()))
473    }
474
475    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
476        self.as_ref().0.poll_flush(cx, false)
477    }
478
479    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
480        self.as_ref().0.poll_shutdown(cx)
481    }
482}
483
484#[cfg(test)]
485mod tests {
486    use std::net;
487
488    use ntex_service::cfg::SharedCfg;
489    use ntex_util::time::{Millis, timeout};
490
491    use super::*;
492
493    #[ntex::test]
494    async fn graceful_shutdown_closes_transport_write_half() {
495        let listener = net::TcpListener::bind("127.0.0.1:0").unwrap();
496        let client = net::TcpStream::connect(listener.local_addr().unwrap()).unwrap();
497        let (server, _) = listener.accept().unwrap();
498        client.set_nonblocking(true).unwrap();
499        server.set_nonblocking(true).unwrap();
500
501        let client = tok_io::net::TcpStream::from_std(client).unwrap();
502        let io = Io::new(
503            super::super::TcpStream(tok_io::net::TcpStream::from_std(server).unwrap()),
504            SharedCfg::default(),
505        );
506
507        io.close();
508
509        let mut buf = [0; 1];
510        let read = timeout(Millis(1000), async {
511            loop {
512                client.readable().await.unwrap();
513                match client.try_read(&mut buf) {
514                    Err(err) if err.kind() == io::ErrorKind::WouldBlock => (),
515                    result => break result,
516                }
517            }
518        })
519        .await
520        .expect("transport write half was not shut down")
521        .unwrap();
522        assert_eq!(read, 0);
523
524        // Keep the Io alive until after EOF is observed. Dropping it must not be
525        // what closes the peer-facing write half.
526        drop(io);
527    }
528
529    struct CaptureCtx(std::rc::Rc<std::cell::Cell<Option<IoContext>>>);
530
531    struct NoHandle;
532
533    impl ntex_io::Handle for NoHandle {}
534
535    impl ntex_io::IoStream for CaptureCtx {
536        fn start(self, ctx: IoContext) -> Box<dyn ntex_io::Handle> {
537            self.0.set(Some(ctx));
538            Box::new(NoHandle)
539        }
540    }
541
542    /// A failed write must hand the pages it took back to the write buffer.
543    /// Dropping them loses the output and leaves it counted as in flight,
544    /// which nothing will ever report as written.
545    #[ntex::test]
546    async fn failed_write_returns_pages_to_the_buffer() {
547        let listener = net::TcpListener::bind("127.0.0.1:0").unwrap();
548        let stream = net::TcpStream::connect(listener.local_addr().unwrap()).unwrap();
549        let _peer = listener.accept().unwrap();
550        stream.set_nonblocking(true).unwrap();
551        // Our own write half is shut, so the next write fails with `EPIPE`.
552        stream.shutdown(net::Shutdown::Write).unwrap();
553        let stream = tok_io::net::TcpStream::from_std(stream).unwrap();
554        // `try_write` reports `WouldBlock` until readiness has been observed.
555        stream.writable().await.unwrap();
556
557        let slot = std::rc::Rc::new(std::cell::Cell::new(None));
558        let io = Io::new(CaptureCtx(slot.clone()), SharedCfg::default());
559        let ctx = slot.take().unwrap();
560        io.encode_slice(b"hello").unwrap();
561
562        let status = write(&stream, &ctx, true);
563
564        assert!(matches!(status, WrtStatus::Terminate));
565        assert_eq!(
566            io.with_write_dst(|b| b.len()),
567            5,
568            "failed write dropped its pages"
569        );
570    }
571}