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 fn terminate(&self) -> io::Result<()>;
35
36 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 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 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 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 let bufs = unsafe { &*(&raw const bufs[..num] as *const [std::io::IoSlice<'_>]) };
286
287 let result = write_io(ctx, io, bufs);
290
291 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 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
347fn 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 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 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 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 #[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 stream.shutdown(net::Shutdown::Write).unwrap();
553 let stream = tok_io::net::TcpStream::from_std(stream).unwrap();
554 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}