use std::{fmt, io, task::Context, task::Poll};
use ntex_bytes::{BytePages, BytesMut};
use ntex_util::time::{Seconds, sleep};
use crate::filter::{read_readiness, write_readiness};
use crate::{Flags, Id, IoRef, IoTaskStatus, Readiness, io::IoState};
#[repr(transparent)]
pub struct IoContext(IoRef);
impl IoContext {
pub(crate) fn new(io: IoRef) -> Self {
Self(io)
}
fn st(&self) -> &IoState {
&self.0.0
}
#[inline]
pub fn id(&self) -> Id {
self.0.id()
}
#[inline]
pub fn tag(&self) -> &'static str {
self.0.tag()
}
#[doc(hidden)]
pub fn flags(&self) -> Flags {
self.0.flags()
}
#[inline]
pub fn shutdown_timeout(&self) -> Seconds {
self.0.cfg().shutdown_timeout()
}
#[inline]
pub fn poll_read_ready(&self, cx: &mut Context<'_>) -> Poll<Readiness> {
let st = self.st();
if st.flags.is_force_closing() {
return Poll::Ready(Readiness::Terminate);
}
self.poll_filters_shutdown(cx);
let res = {
let _borrow = st.buffer.borrow();
self.0.filter().poll_read_ready(cx)
};
if res.is_pending() {
if !st.flags.is_read_filter_paused()
&& read_readiness(st) == Poll::Ready(Readiness::Ready)
{
log::trace!("{}: Filter is not ready, pause reading", st.tag());
st.flags.set_read_filter_paused();
st.wake_dispatch_task();
}
} else if st.flags.is_read_filter_paused() {
log::trace!("{}: Filter is ready, resume reading", st.tag());
st.flags.unset_read_filter_paused();
st.wake_dispatch_task();
}
res
}
#[inline]
pub fn poll_write_ready(&self, cx: &mut Context<'_>) -> Poll<Readiness> {
let st = self.st();
if st.flags.is_force_closing() {
return Poll::Ready(Readiness::Terminate);
}
self.poll_shutdown_deadline(cx);
let res = {
let _borrow = st.buffer.borrow();
self.0.filter().poll_write_ready(cx)
};
if res.is_pending() {
if !st.flags.is_write_filter_paused()
&& write_readiness(st) == Poll::Ready(Readiness::Ready)
{
log::trace!("{}: Filter is not ready, pause writing", st.tag());
st.flags.set_write_filter_paused();
st.wake_dispatch_task();
}
} else if st.flags.is_write_filter_paused() {
log::trace!("{}: Filter is ready, resume writing", st.tag());
st.flags.unset_write_filter_paused();
st.wake_dispatch_task();
}
res
}
pub fn stop(&self, e: Option<io::Error>) {
self.st().terminate_connection(e);
}
pub fn stopped(&self, e: Option<io::Error>) {
self.st().stop_connection(e);
}
pub fn take_read_buf(&self) -> BytesMut {
let st = self.st();
if st.flags.is_read_ready() {
st.get_read_buf()
} else if let Some(mut buf) = st.buffer.get_read_buf() {
buf.reserve_more();
buf
} else {
st.get_read_buf()
}
}
pub fn release_read_buf(&self, buf: BytesMut, status: Poll<io::Result<usize>>) -> IoTaskStatus {
let st = self.st();
let orig = st.buffer.read_dst_size();
#[cfg(feature = "trace")]
log::trace!(
"{}: read-status == {status:?} orig:{orig:?} flags:{:?}",
st.tag(),
st.flags
);
if st.flags.is_stopping() {
let mut buf = buf;
buf.clear();
st.buffer.set_read_buf(buf);
stopping_read_status(st, &status)
} else {
let mut buf = buf;
track_read(st, &buf, &status);
if st.is_io_dropped() {
buf.clear();
}
st.buffer.set_read_buf(buf);
self.process_read_status(orig, status)
}
}
pub fn with_read_buf<F>(&self, f: F) -> IoTaskStatus
where
F: FnOnce(&mut BytesMut) -> Poll<io::Result<usize>>,
{
let st = self.st();
let orig = st.buffer.read_dst_size();
let stopping = st.flags.is_stopping();
let discard = stopping || st.is_io_dropped();
let status = st.buffer.with_read_src(&self.0, |buf| {
buf.reserve_more();
let status = f(buf);
if !stopping {
track_read(st, buf, &status);
}
if discard {
buf.clear();
}
status
});
#[cfg(feature = "trace")]
log::trace!(
"{}: rd-status = {status:?} orig:{orig:?} flags:{:?}",
st.tag(),
st.flags
);
if stopping {
stopping_read_status(st, &status)
} else {
self.process_read_status(orig, status)
}
}
fn process_read_status(&self, orig: usize, status: Poll<io::Result<usize>>) -> IoTaskStatus {
let st = self.st();
let result = match status {
Poll::Pending => Ok(()),
Poll::Ready(status) => status.and_then(|nbytes| {
if nbytes == 0 {
if st.flags.is_read_eof() {
return Ok(());
}
st.flags.set_read_eof();
st.wake_dispatch_task();
}
st.buffer.process_read_buf(&self.0).and_then(|status| {
let size = st.buffer.read_dst_size();
if size > orig {
if st.is_rd_backpressure_needed(size) {
log::trace!("{}: Read buf({size}), enable back-pressure", st.tag());
st.flags.set_read_ready_and_backpressure();
} else {
st.flags.set_read_ready();
}
#[cfg(feature = "trace")]
log::trace!("{}: New {size} bytes available", st.tag());
st.wake_dispatch_task();
}
if st.flags.is_read_notify() {
st.wake_dispatch_task();
st.flags.set_read_notified();
}
if status.wants_write {
st.buffer.process_write_buf_force(&self.0)?;
self.0.consolidate_write_state(false)?;
if st.is_wr_backpressure_needed(st.transport_outstanding()) {
log::trace!("{}: Write buf is full, pause reading", st.tag());
st.flags.set_read_wr_backpressure();
}
}
if st.flags.is_shutting_down_filters() {
st.wake_read_task();
}
Ok(())
})
}),
};
if let Err(err) = result {
if st.flags.is_stopping_filters() {
let _ = st.buffer.process_write_buf_force(&self.0);
stop_filters(st, Some(err));
IoTaskStatus::Pause
} else {
st.terminate_connection(Some(err));
IoTaskStatus::Stop
}
} else if st.flags.is_aborted() {
IoTaskStatus::Stop
} else if st.flags.is_read_eof()
|| st.flags.is_read_paused_or_backpressure()
|| st.flags.is_read_filter_paused()
|| (st.flags.is_read_wr_backpressure() && !st.flags.is_stopping_filters())
{
IoTaskStatus::Pause
} else {
IoTaskStatus::Io
}
}
pub fn with_write_dst<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut BytePages) -> R,
{
let st = self.st();
if let Err(e) = st.buffer.process_write_buf(&self.0) {
st.terminate_connection(Some(e));
}
let before = st.buffer.write_buf_size();
let result = st.buffer.with_write_dst(|buffer| f(buffer));
st.track_wr_inflight(before, st.buffer.write_buf_size());
result
}
pub fn update_write_status(&self, status: io::Result<usize>) -> IoTaskStatus {
let st = &self.st();
#[cfg(feature = "trace")]
log::trace!(
"{}: write-status == {status:?} buf:{} inflight:{} flags:{:?}",
st.tag(),
st.buffer.write_buf_size(),
st.wr_inflight.get(),
st.flags
);
match status {
Ok(written) => {
st.wr_inflight_written(written);
let len = st.buffer.write_buf_size();
let outstanding = st.write_outstanding();
if st.flags.is_write_flush() {
if outstanding == 0 {
st.wake_dispatch_task();
}
} else if st.flags.is_wr_backpressure()
&& st.should_disable_wr_backpressure(outstanding)
{
st.wake_dispatch_task();
}
if st.flags.is_wr_backpressure() && st.should_disable_wr_backpressure(outstanding) {
st.wake_write_waiters();
}
if st.flags.is_read_wr_backpressure()
&& st.should_disable_wr_backpressure(st.transport_outstanding())
{
st.flags.unset_read_wr_backpressure();
st.wake_read_task();
}
if st.flags.is_aborted() {
IoTaskStatus::Stop
} else if len == 0 {
st.flags.set_write_paused();
if st.flags.is_stopping_filters() {
st.wake_read_task();
}
if st.flags.is_stopping() && outstanding == 0 {
st.wake_write_task();
}
IoTaskStatus::Pause
} else {
st.flags.unset_write_paused();
if st.flags.is_write_filter_paused() {
IoTaskStatus::Pause
} else {
IoTaskStatus::Io
}
}
}
Err(err) => {
st.terminate_connection(Some(err));
IoTaskStatus::Stop
}
}
}
fn poll_filters_shutdown(&self, cx: &mut Context<'_>) {
let st = &self.st();
if !st.flags.is_shutting_down_filters() {
return;
}
let ready = match st.buffer.process_shutdown(&self.0) {
Ok(Poll::Ready(())) => true,
Ok(Poll::Pending) => false,
Err(err) => {
st.terminate_connection(Some(err));
return;
}
};
if self.0.consolidate_write_state(true).is_err() {
return;
}
let flushed = st.flags.is_write_paused() && !st.flags.is_wr_send_scheduled();
#[cfg(feature = "trace")]
log::trace!(
"{}: shutdown filters, done:{ready:?} flushed:{flushed:?} wr-buf:{:?}, flags:{:?}",
st.tag(),
st.buffer.write_buf_size(),
st.flags,
);
if ready && flushed {
st.filters_stopped();
return;
}
let eof = !ready && st.flags.is_read_eof();
let blocked =
!ready && !eof && (st.flags.is_read_paused() || st.flags.is_rd_backpressure());
if eof || blocked {
if eof {
log::debug!("{}: Peer closed before filter shutdown completed", st.tag());
}
stop_filters(st, blocked.then(blocked_err));
return;
}
let timeout = st
.shutdown_timeout
.take()
.unwrap_or_else(|| sleep(st.cfg.shutdown_timeout()));
if timeout.poll_elapsed(cx).is_ready() {
stop_filters(
st,
Some(io::Error::new(
io::ErrorKind::TimedOut,
"filter shutdown timed out",
)),
);
}
st.shutdown_timeout.set(Some(timeout));
}
fn poll_shutdown_deadline(&self, cx: &mut Context<'_>) {
let st = &self.st();
if !st.flags.is_stopping() {
return;
}
if st.write_outstanding() == 0 {
return;
}
let timeout = st
.shutdown_timeout
.take()
.unwrap_or_else(|| sleep(st.cfg.shutdown_timeout()));
if timeout.poll_elapsed(cx).is_ready() {
let len = st.write_outstanding();
if len != 0 {
log::warn!(
"{}: Shutdown timed out, discarding {len} bytes of buffered output",
st.tag()
);
}
st.terminate_connection(Some(io::Error::new(
io::ErrorKind::TimedOut,
"io shutdown timed out",
)));
} else {
st.shutdown_timeout.set(Some(timeout));
}
}
}
fn blocked_err() -> io::Error {
io::Error::other("filter shutdown blocked by unread buffered data")
}
fn stop_filters(st: &IoState, err: Option<io::Error>) {
if let Some(err) = err {
st.set_shutdown_error(err);
}
st.filters_stopped();
}
fn track_read(st: &IoState, buf: &BytesMut, status: &Poll<io::Result<usize>>) {
match status {
Poll::Ready(Ok(n)) => st.track_read(*n, *n != 0 && buf.len() == buf.capacity()),
Poll::Pending => st.track_read(0, false),
Poll::Ready(Err(_)) => {}
}
}
fn stopping_read_status(st: &IoState, status: &Poll<io::Result<usize>>) -> IoTaskStatus {
match status {
Poll::Ready(Ok(n)) if *n != 0 => IoTaskStatus::Io,
Poll::Ready(_) => {
st.flags.set_read_eof();
IoTaskStatus::Pause
}
Poll::Pending => IoTaskStatus::Pause,
}
}
impl Clone for IoContext {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl fmt::Debug for IoContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("IoContext").field("io", &self.0).finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{FilterBuf, FilterLayer, Io, testing::IoTest};
use ntex_util::future::lazy;
#[ntex::test]
async fn ctx_basics() {
let (_, server) = IoTest::create();
let state = Io::from(server);
let ctx = IoContext::new(state.get_ref());
let _ = ctx.flags();
assert_ne!(ctx.id(), Id::default());
assert!(format!("{ctx:?}").contains("IoContext"));
}
#[ntex::test]
async fn pending_read_completion_is_not_eof() {
let (_, server) = IoTest::create();
let state = Io::from(server);
let ctx = IoContext::new(state.get_ref());
assert!(lazy(|cx| state.poll_read_more(cx)).await.is_pending());
assert_ne!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Pending),
IoTaskStatus::Stop
);
assert!(lazy(|cx| state.poll_read_more(cx)).await.is_pending());
assert_eq!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Ready(Ok(0))),
IoTaskStatus::Pause
);
assert!(matches!(
lazy(|cx| state.poll_read_more(cx)).await,
Poll::Ready(Ok(None))
));
}
#[derive(Debug)]
struct FinishOnEof;
impl FilterLayer for FinishOnEof {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
if buf.io().is_read_eof() {
buf.with_read_buffers(|_, dst| dst.extend_from_slice(b"final"));
}
Ok(())
}
fn process_write_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
}
#[ntex::test]
async fn eof_reports_available_input_once() {
let (_, server) = IoTest::create();
let state = Io::from(server);
let ctx = IoContext::new(state.get_ref());
ctx.release_read_buf(BytesMut::copy_from_slice(b"12345"), Poll::Ready(Ok(5)));
ctx.release_read_buf(ctx.take_read_buf(), Poll::Ready(Ok(0)));
assert!(matches!(
lazy(|cx| state.poll_read_more(cx)).await,
Poll::Ready(Ok(Some(())))
));
assert_eq!(state.with_read_dst(|b| b.len()), 5);
assert!(matches!(
lazy(|cx| state.poll_read_more(cx)).await,
Poll::Ready(Ok(None))
));
assert_eq!(state.with_read_dst(BytesMut::take), b"12345");
}
#[ntex::test]
async fn shutdown_keeps_unconsumed_input_visible() {
let (_, server) = IoTest::create();
let state = Io::from(server);
let ctx = IoContext::new(state.get_ref());
ctx.release_read_buf(BytesMut::copy_from_slice(b"12345"), Poll::Ready(Ok(5)));
assert!(ctx.flags().is_read_ready());
assert!(lazy(|cx| state.poll_shutdown(cx)).await.is_pending());
assert!(ctx.flags().is_read_ready());
assert!(ctx.take_read_buf().is_empty());
assert_eq!(state.with_read_dst(BytesMut::take), b"12345");
}
#[ntex::test]
async fn clean_eof_is_processed_by_filters_once() {
let (_, server) = IoTest::create();
let state = Io::from(server).add_filter(FinishOnEof);
let ctx = IoContext::new(state.get_ref());
for _ in 0..3 {
assert_eq!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Ready(Ok(0))),
IoTaskStatus::Pause
);
assert!(state.is_read_eof());
}
assert_eq!(state.with_read_dst(BytesMut::take), b"final");
}
#[ntex::test]
async fn clean_eof_is_processed_by_filters() {
let (_, server) = IoTest::create();
let state = Io::from(server).add_filter(FinishOnEof);
let ctx = IoContext::new(state.get_ref());
assert!(lazy(|cx| state.poll_read_notify(cx)).await.is_pending());
assert_eq!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Ready(Ok(0))),
IoTaskStatus::Pause
);
assert!(state.is_read_eof());
assert!(matches!(
lazy(|cx| state.poll_read_notify(cx)).await,
Poll::Ready(Ok(Some(())))
));
assert!(matches!(
lazy(|cx| state.poll_read_notify(cx)).await,
Poll::Ready(Ok(None))
));
assert_eq!(state.with_read_dst(BytesMut::take), b"final");
assert!(matches!(
lazy(|cx| state.poll_read_more(cx)).await,
Poll::Ready(Ok(None))
));
}
#[derive(Debug)]
struct RejectEof;
impl FilterLayer for RejectEof {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
if buf.io().is_read_eof() {
Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"truncated filtered stream",
))
} else {
Ok(())
}
}
fn process_write_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
}
#[ntex::test]
async fn clean_eof_filter_error_terminates_connection() {
let (_, server) = IoTest::create();
let state = Io::from(server).add_filter(RejectEof);
let ctx = IoContext::new(state.get_ref());
assert_eq!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Ready(Ok(0))),
IoTaskStatus::Stop
);
assert!(state.is_read_eof());
assert!(state.flags().is_terminating());
}
#[ntex::test]
async fn filter_shutdown_does_not_spin_on_pending_read() {
use crate::{Handle, IoStream};
use std::{cell::Cell, future::poll_fn, rc::Rc};
struct Stalled(Rc<Cell<usize>>);
impl IoStream for Stalled {
fn start(self, ctx: IoContext) -> Box<dyn Handle> {
let polls = self.0.clone();
ntex_util::spawn(async move {
poll_fn(|cx| {
polls.set(polls.get() + 1);
if let Poll::Ready(Readiness::Ready) = ctx.poll_read_ready(cx) {
let _ = ctx.with_read_buf(|_| Poll::Pending);
}
match ctx.poll_write_ready(cx) {
Poll::Ready(Readiness::Ready) => {
let _ = ctx.update_write_status(Ok(0));
Poll::Pending
}
Poll::Ready(_) => Poll::Ready(()),
Poll::Pending => Poll::Pending,
}
})
.await;
ctx.stopped(None);
});
Box::new(Stalled(self.0))
}
}
impl Handle for Stalled {}
let polls = Rc::new(Cell::new(0));
let io = Io::new(
Stalled(polls.clone()),
ntex_service::cfg::SharedCfg::default(),
);
io.encode_slice(b"data").unwrap();
ntex_util::time::sleep(ntex_util::time::Millis(20)).await;
io.close();
let start = polls.get();
ntex_util::time::sleep(ntex_util::time::Millis(100)).await;
assert!(
polls.get() - start < 10,
"io task polled {} times",
polls.get() - start
);
}
struct Manual;
impl crate::IoStream for Manual {
fn start(self, _: IoContext) -> Box<dyn crate::Handle> {
Box::new(Manual)
}
}
impl crate::Handle for Manual {}
#[ntex::test]
async fn take_read_buf_reuses_consumed_buffer() {
let io = Io::new(Manual, ntex_service::cfg::SharedCfg::default());
let ctx = IoContext::new(io.get_ref());
assert_eq!(ctx.shutdown_timeout(), io.cfg().shutdown_timeout());
assert!(io.query::<u32>().get().is_none());
let mut buf = ctx.take_read_buf();
buf.extend_from_slice(b"12345");
assert_eq!(
ctx.release_read_buf(buf, Poll::Ready(Ok(5))),
IoTaskStatus::Io
);
assert_eq!(io.with_read_dst(|b| b.split_to(3)), b"123");
let mut buf = ctx.take_read_buf();
assert_eq!(buf, b"45");
buf.reserve_more();
assert!(buf.capacity() - buf.len() >= io.get_ref().0.read_size().low());
buf.extend_from_slice(b"6");
assert_eq!(
ctx.release_read_buf(buf, Poll::Ready(Ok(1))),
IoTaskStatus::Io
);
assert_eq!(io.with_read_dst(BytesMut::take), b"456");
}
#[ntex::test]
async fn reads_are_discarded_in_transport_shutdown_phase() {
let io = Io::new(Manual, ntex_service::cfg::SharedCfg::default());
let ctx = IoContext::new(io.get_ref());
io.close();
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(ctx.flags().is_stopping());
let mut buf = ctx.take_read_buf();
buf.extend_from_slice(b"12345");
assert_eq!(
ctx.release_read_buf(buf, Poll::Ready(Ok(5))),
IoTaskStatus::Io
);
assert_eq!(
ctx.with_read_buf(|buf| {
buf.extend_from_slice(b"678");
Poll::Ready(Ok(3))
}),
IoTaskStatus::Io
);
assert_eq!(io.with_read_dst(|b| b.len()), 0);
assert_eq!(
ctx.release_read_buf(ctx.take_read_buf(), Poll::Pending),
IoTaskStatus::Pause
);
assert_eq!(ctx.with_read_buf(|_| Poll::Pending), IoTaskStatus::Pause);
assert!(!io.is_read_eof());
assert_eq!(
ctx.release_read_buf(
ctx.take_read_buf(),
Poll::Ready(Err(io::Error::other("err")))
),
IoTaskStatus::Pause
);
assert!(io.is_read_eof());
assert!(!ctx.flags().is_terminating());
assert_eq!(
ctx.with_read_buf(|_| Poll::Ready(Ok(0))),
IoTaskStatus::Pause
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
ctx.stopped(None);
assert!(io.is_closed());
}
#[derive(Debug)]
struct FailShutdown;
impl FilterLayer for FailShutdown {
fn process_read_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn process_write_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Err(io::Error::other("shutdown failed"))
}
}
#[ntex::test]
async fn filter_shutdown_error_terminates_connection() {
let io = Io::new(Manual, ntex_service::cfg::SharedCfg::default()).add_filter(FailShutdown);
let ctx = IoContext::new(io.get_ref());
assert!(io.query::<u32>().get().is_none());
io.close();
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
assert!(ctx.flags().is_terminating());
ctx.stopped(None);
let err = io.shutdown().await.unwrap_err();
assert_eq!(err.to_string(), "shutdown failed");
}
#[ntex::test]
async fn read_after_termination_stops_read_task() {
let io = Io::new(Manual, ntex_service::cfg::SharedCfg::default());
let ctx = IoContext::new(io.get_ref());
let ctx2 = ctx.clone();
assert_eq!(ctx.id(), ctx2.id());
ctx.stop(Some(io::Error::other("failed")));
let mut buf = ctx2.take_read_buf();
buf.extend_from_slice(b"1");
assert_eq!(
ctx2.release_read_buf(buf, Poll::Ready(Ok(1))),
IoTaskStatus::Stop
);
}
}