use std::cell::{Cell, UnsafeCell};
use std::task::{Context, Poll};
use std::{fmt, future::poll_fn, hash, io, marker, mem, ops, ptr, rc::Rc};
use ntex_bytes::{BytePageSize, BytesMut};
use ntex_codec::{Decoder, Encoder};
use ntex_service::cfg::{Cfg, SharedCfg};
use ntex_util::{future::Either, task::LocalWaker, time::Sleep};
use crate::buf::Stack;
use crate::cfg::IoConfig;
use crate::ctx::IoContext;
use crate::filter::{Base, Filter, Layer};
use crate::filterptr::FilterPtr;
use crate::flags::Flags;
use crate::ops::{Id, IoManager, TimerHandle};
use crate::seal::{IoBoxed, Sealed};
use crate::utils::{Extensions, WriteDeadline, write_timed_out};
use crate::waiters::{TAG_WRITE, WriteGuard};
use crate::{Decoded, FilterLayer, Handle, IoStatusUpdate, IoStream, RecvError};
pub struct Io<F = Base>(UnsafeCell<IoRef>, marker::PhantomData<F>);
#[derive(Clone)]
pub struct IoRef(pub(super) Rc<IoState>);
pub(crate) struct IoState {
filter: FilterPtr,
pub(super) id: Cell<Id>,
pub(super) cfg: Cfg<IoConfig>,
pub(super) flags: Flags,
pub(super) error: Cell<Option<io::Error>>,
pub(super) read_task: LocalWaker,
pub(super) write_task: LocalWaker,
dispatch_task: LocalWaker,
pub(super) buffer: Stack,
pub(super) handle: Cell<Option<Box<dyn Handle>>>,
pub(super) timeout: Cell<TimerHandle>,
pub(super) shutdown_timeout: Cell<Option<Sleep>>,
pub(super) wr_inflight: Cell<u32>,
pub(super) extensions: Extensions,
}
impl IoState {
pub(super) fn id(&self) -> Id {
self.id.get()
}
pub(super) fn tag(&self) -> &'static str {
self.cfg.tag()
}
pub(super) fn is_io_dropped(&self) -> bool {
!self.filter.is_set()
}
pub(super) fn filter(&self) -> &dyn Filter {
self.filter.get()
}
pub(super) fn notify_timeout(&self) {
if self.flags.check_dispatcher_timeout_unset() {
self.wake_dispatch_task();
log::trace!("{}: Timer, notify dispatcher", self.cfg.tag());
}
}
pub(super) fn notify_disconnect(&self) {
self.extensions.wake_all();
}
pub(super) fn error(&self) -> Option<io::Error> {
if let Some(err) = self.error.take() {
let cloned = if let Some(code) = err.raw_os_error() {
io::Error::from_raw_os_error(code)
} else {
io::Error::new(err.kind(), format!("{err}"))
};
self.error.set(Some(cloned));
Some(err)
} else {
None
}
}
pub(super) fn error_or_disconnected(&self) -> io::Error {
self.error()
.unwrap_or_else(|| io::Error::new(io::ErrorKind::NotConnected, "Disconnected"))
}
pub(super) fn filters_stopped(&self) {
self.wake_read_task();
self.wake_write_task();
self.wake_dispatch_task();
self.wake_write_waiters();
self.flags.enter_transport_shutdown();
}
fn set_error(&self, err: Option<io::Error>) {
if let Some(err) = err {
if let Some(current) = self.error.take() {
self.error.set(Some(current));
} else {
self.error.set(Some(err));
}
}
}
pub(super) fn set_shutdown_error(&self, err: io::Error) {
self.set_error(Some(err));
}
pub(super) fn force_close_connection(&self) {
self.begin_terminate(None, true);
}
pub(super) fn terminate_connection(&self, err: Option<io::Error>) {
self.begin_terminate(err, false);
}
fn begin_terminate(&self, err: Option<io::Error>, force: bool) {
self.set_error(err);
if self.flags.begin_terminate(force) {
log::trace!("{}: Terminate io", self.cfg.tag());
self.wr_inflight.set(0);
self.wake_read_task();
self.wake_write_task();
self.wake_dispatch_task();
self.wake_write_waiters();
self.handle.take();
}
}
pub(super) fn stop_connection(&self, err: Option<io::Error>) {
if !self.flags.is_closed() {
log::trace!("{}: Stop io with error {:?}", self.cfg.tag(), err);
self.set_error(err);
self.flags.set_stopped();
self.wr_inflight.set(0);
self.wake_read_task();
self.wake_write_task();
self.wake_dispatch_task();
self.wake_write_waiters();
self.notify_disconnect();
self.handle.take();
}
}
pub(super) fn start_shutdown(&self) {
if self.flags.is_active() {
log::trace!("{}: Initiate io shutdown {:?}", self.cfg.tag(), self.flags);
self.flags.enter_filters_stopping();
self.wake_read_task();
self.wake_write_task();
}
}
pub(super) fn get_read_buf(&self) -> BytesMut {
self.cfg.read_buf().get()
}
pub(super) fn is_rd_backpressure_needed(&self, size: usize) -> bool {
size >= self.cfg.read_buf().high
}
pub(super) fn is_wr_backpressure_needed(&self, size: usize) -> bool {
size >= self.cfg.write_buf().high
}
pub(super) fn should_disable_rd_backpressure(&self, size: usize) -> bool {
size <= self.cfg.read_buf().half
}
pub(super) fn should_disable_wr_backpressure(&self, size: usize) -> bool {
size <= self.cfg.write_buf().half
}
pub(super) fn write_outstanding(&self) -> usize {
self.buffer.write_buf_size() + self.wr_inflight.get() as usize
}
pub(super) fn transport_outstanding(&self) -> usize {
self.buffer.write_dst_size() + self.wr_inflight.get() as usize
}
pub(super) fn track_wr_inflight(&self, before: usize, after: usize) {
let inflight = self.wr_inflight.get();
if after < before {
self.wr_inflight
.set(inflight.saturating_add(as_u32(before - after)));
} else {
self.wr_inflight
.set(inflight.saturating_sub(as_u32(after - before)));
}
}
pub(super) fn wr_inflight_written(&self, written: usize) {
self.wr_inflight
.set(self.wr_inflight.get().saturating_sub(as_u32(written)));
}
pub(super) fn wake_read_task(&self) {
self.read_task.wake();
}
pub(super) fn wake_write_task(&self) {
#[cfg(feature = "trace")]
log::trace!("{}: Wake write task, flags:{:?}", self.tag(), self.flags);
self.write_task.wake();
}
pub(super) fn wake_dispatch_task(&self) {
self.dispatch_task.wake();
}
pub(super) fn wake_write_waiters(&self) {
self.extensions.wake(TAG_WRITE);
}
pub(super) fn check_write_ready(&self) -> Option<io::Result<()>> {
if self.flags.is_peer_gone() {
Some(Err(self.error_or_disconnected()))
} else if !self.flags.is_wr_backpressure()
|| self.should_disable_wr_backpressure(self.write_outstanding())
{
Some(Ok(()))
} else {
None
}
}
pub(super) async fn write_ready(&self) -> io::Result<()> {
if let Some(res) = self.check_write_ready() {
return res;
}
let waiter = WriteGuard::new(&self.extensions);
let mut deadline = WriteDeadline::new(self.cfg.write_timeout());
poll_fn(|cx| {
if let Some(res) = self.check_write_ready() {
Poll::Ready(res)
} else if deadline.poll_expired(cx) {
Poll::Ready(Err(write_timed_out()))
} else {
waiter.register(cx);
Poll::Pending
}
})
.await
}
pub(super) async fn with_write_timeout<T, F>(&self, mut f: F) -> io::Result<T>
where
F: FnMut(&mut Context<'_>) -> Poll<io::Result<T>>,
{
let mut deadline = WriteDeadline::new(self.cfg.write_timeout());
poll_fn(|cx| match f(cx) {
Poll::Ready(res) => Poll::Ready(res),
Poll::Pending if deadline.poll_expired(cx) => Poll::Ready(Err(write_timed_out())),
Poll::Pending => Poll::Pending,
})
.await
}
}
impl Eq for IoState {}
impl PartialEq for IoState {
#[inline]
fn eq(&self, other: &Self) -> bool {
ptr::eq(self, other)
}
}
impl hash::Hash for IoState {
#[inline]
fn hash<H: hash::Hasher>(&self, state: &mut H) {
(ptr::from_ref(self) as usize).hash(state);
}
}
impl fmt::Debug for IoState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let err = self.error.take();
let res = f
.debug_struct("IoState")
.field("id", &self.id)
.field("flags", &self.flags)
.field("filter", &self.filter.is_set())
.field("timeout", &self.timeout)
.field("error", &err)
.field("buffer", &self.buffer)
.field("cfg", &self.cfg)
.finish();
self.error.set(err);
res
}
}
impl Io {
pub fn new<I: IoStream, T: Into<SharedCfg>>(io: I, cfg: T) -> Self {
let cfg = cfg.into().get::<IoConfig>();
let size = cfg.write_page_size();
let flags = Flags::new(cfg.write_buf_threshold() > 0);
let inner = Rc::new(IoState {
cfg,
flags,
id: Cell::new(Id::default()),
filter: FilterPtr::null(),
error: Cell::new(None),
dispatch_task: LocalWaker::new(),
read_task: LocalWaker::new(),
write_task: LocalWaker::new(),
buffer: Stack::new(size),
handle: Cell::new(None),
timeout: Cell::new(TimerHandle::default()),
shutdown_timeout: Cell::new(None),
wr_inflight: Cell::new(0),
extensions: Extensions::default(),
});
inner.filter.set(Base::new(IoRef(inner.clone())));
let ioref = IoRef(inner);
ioref.0.id.set(IoManager::register(&ioref));
let hnd = io.start(IoContext::new(ioref.clone()));
ioref.0.handle.set(Some(hnd));
Io(UnsafeCell::new(ioref), marker::PhantomData)
}
}
impl<I: IoStream> From<I> for Io {
#[inline]
fn from(io: I) -> Io {
Io::new(io, SharedCfg::default())
}
}
impl IoRef {
fn create_empty() -> IoRef {
IoRef(Rc::new(IoState {
id: Cell::new(Id::default()),
cfg: SharedCfg::default().get::<IoConfig>(),
filter: FilterPtr::null(),
flags: Flags::new_stopped(),
error: Cell::new(None),
dispatch_task: LocalWaker::new(),
read_task: LocalWaker::new(),
write_task: LocalWaker::new(),
buffer: Stack::new(BytePageSize::Size16),
handle: Cell::new(None),
timeout: Cell::new(TimerHandle::default()),
shutdown_timeout: Cell::new(None),
wr_inflight: Cell::new(0),
extensions: Extensions::default(),
}))
}
}
impl<F> Io<F> {
#[inline]
pub fn get_ref(&self) -> IoRef {
self.io_ref().clone()
}
#[inline]
#[must_use]
pub unsafe fn take(&self) -> Self {
Self(UnsafeCell::new(self.take_io_ref()), marker::PhantomData)
}
fn take_io_ref(&self) -> IoRef {
unsafe { mem::replace(&mut *self.0.get(), IoRef::create_empty()) }
}
#[track_caller]
fn check_not_borrowed(&self) {
if self.st().buffer.is_borrowed() {
let tag = self.tag();
mem::forget(self.take_io_ref());
panic!("{tag}: filter chain is changed while it is in use");
}
}
fn st(&self) -> &IoState {
unsafe { &(*self.0.get()).0 }
}
fn io_ref(&self) -> &IoRef {
unsafe { &*self.0.get() }
}
#[inline]
pub unsafe fn set_config<T: Into<SharedCfg>>(&self, cfg: T) {
let cfg = cfg.into().get::<IoConfig>();
let page_size = cfg.write_page_size();
if self.cfg().write_page_size() != page_size {
self.st().buffer.set_page_size(page_size);
}
self.st()
.flags
.set_direct_wr_enabled(cfg.write_buf_threshold() > 0);
unsafe {
self.st().cfg.replace(cfg);
}
}
}
impl<F: FilterLayer, T: Filter> Io<Layer<F, T>> {
#[inline]
pub fn filter(&self) -> &F {
&self.st().filter.filter::<Layer<F, T>>().0
}
}
impl<F: Filter> Io<F> {
#[inline]
pub fn seal(self) -> Io<Sealed> {
self.check_not_borrowed();
let state = self.take_io_ref();
state.0.filter.seal::<F>();
Io(UnsafeCell::new(state), marker::PhantomData)
}
#[inline]
pub fn boxed(self) -> IoBoxed {
self.seal().into()
}
#[inline]
pub fn add_filter<U>(self, nf: U) -> Io<Layer<U, F>>
where
U: FilterLayer,
{
self.check_not_borrowed();
self.with_callbacks(|cb| cb.before_processing(&self));
if let Err(e) = self.st().buffer.process_write_buf_no_cb(&self) {
self.st().terminate_connection(Some(e));
}
let state = self.take_io_ref();
state.0.buffer.add_layer(state.0.cfg.write_page_size());
state.0.filter.add_filter::<F, U>(nf);
let io = Io(UnsafeCell::new(state), marker::PhantomData);
if let Err(e) = io.st().buffer.process_read_buf_no_cb(&io) {
io.st().terminate_connection(Some(e));
}
io.with_callbacks(|cb| cb.after_processing(&io));
io
}
#[allow(clippy::items_after_statements)]
pub fn map_filter<U, R>(self, f: U) -> Io<R>
where
U: FnOnce(F) -> R,
R: Filter,
{
self.check_not_borrowed();
self.with_callbacks(|cb| cb.before_processing(&self));
if let Err(e) = self.st().buffer.process_write_buf(&self) {
self.st().terminate_connection(Some(e));
}
struct Guard<'a>(&'a IoRef);
impl Drop for Guard<'_> {
fn drop(&mut self) {
let st = &self.0.0;
st.force_close_connection();
st.buffer.release(self.0.cfg());
drop(st.extensions.take_callbacks());
}
}
let guard = Guard(self.io_ref());
self.st().filter.map_filter::<F, U, R>(f);
mem::forget(guard);
let state = self.take_io_ref();
let io = Io(UnsafeCell::new(state), marker::PhantomData);
io.with_callbacks(|cb| cb.after_processing(&io));
io
}
}
impl<F> Io<F> {
pub async fn recv<U>(&self, codec: &U) -> Result<Option<U::Item>, Either<U::Error, io::Error>>
where
U: Decoder,
{
loop {
return match poll_fn(|cx| self.poll_recv(codec, cx)).await {
Ok(item) => Ok(Some(item)),
Err(RecvError::Timeout) => Err(Either::Right(io::Error::new(
io::ErrorKind::TimedOut,
"Timeout",
))),
Err(RecvError::WriteBackpressure) => {
let timed_out = poll_fn(|cx| {
if self.st().flags.check_dispatcher_timeout() {
Poll::Ready(Ok(true))
} else {
self.poll_flush(cx, false).map_ok(|()| false)
}
})
.await
.map_err(Either::Right)?;
if timed_out {
Err(Either::Right(io::Error::new(
io::ErrorKind::TimedOut,
"Timeout",
)))
} else {
continue;
}
}
Err(RecvError::Decoder(err)) => Err(Either::Left(err)),
Err(RecvError::PeerGone(Some(err))) => Err(Either::Right(err)),
Err(RecvError::PeerGone(None)) => {
let st = self.st();
if st.flags.is_read_eof() && st.buffer.read_dst_size() != 0 {
Err(Either::Right(io::Error::new(
io::ErrorKind::UnexpectedEof,
"bytes remaining on stream",
)))
} else {
Ok(None)
}
}
};
}
}
pub async fn read_exact(&self, dst: &mut [u8]) -> io::Result<()> {
loop {
let completed = self.with_read_dst(|buf| {
if buf.len() >= dst.len() {
let _ = io::Read::read(buf, dst).expect("Cannot fail");
true
} else {
false
}
});
if completed {
return Ok(());
}
if self.read_more().await?.is_none() {
return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "Disconnected"));
}
}
}
#[inline]
pub async fn read_more(&self) -> io::Result<Option<()>> {
poll_fn(|cx| self.poll_read_more(cx)).await
}
#[inline]
pub async fn read_notify(&self) -> io::Result<Option<()>> {
poll_fn(|cx| self.poll_read_notify(cx)).await
}
#[inline]
pub async fn send<U>(&self, item: U::Item, codec: &U) -> Result<(), Either<U::Error, io::Error>>
where
U: Encoder,
{
self.encode(item, codec).map_err(Either::Left)?;
self.st()
.with_write_timeout(|cx| self.poll_flush(cx, true))
.await
.map_err(Either::Right)?;
Ok(())
}
#[inline]
pub async fn flush(&self, full: bool) -> io::Result<()> {
self.st()
.with_write_timeout(|cx| self.poll_flush(cx, full))
.await
}
#[inline]
pub async fn shutdown(&self) -> io::Result<()> {
poll_fn(|cx| self.poll_shutdown(cx)).await
}
#[inline]
pub fn poll_read_more(&self, cx: &mut Context<'_>) -> Poll<io::Result<Option<()>>> {
let st = self.st();
if st.flags.is_peer_gone() {
if let Some(err) = st.error() {
Poll::Ready(Err(err))
} else {
Poll::Ready(Ok(None))
}
} else {
let ready = st.flags.is_read_ready();
if st.flags.is_read_eof() && !ready {
return Poll::Ready(Ok(None));
}
if st.flags.is_read_paused_or_backpressure() || st.flags.is_read_wr_backpressure() {
st.flags.unset_read_ready_and_backpressure();
st.flags.unset_read_paused();
st.flags.unset_read_wr_backpressure();
st.wake_read_task();
if ready {
Poll::Ready(Ok(Some(())))
} else {
st.dispatch_task.register(cx.waker());
Poll::Pending
}
} else if ready {
Poll::Ready(Ok(Some(())))
} else {
st.dispatch_task.register(cx.waker());
Poll::Pending
}
}
}
#[inline]
pub fn poll_read_notify(&self, cx: &mut Context<'_>) -> Poll<io::Result<Option<()>>> {
let st = self.st();
if st.flags.is_stopping_or_terminating() {
if let Some(err) = st.error() {
Poll::Ready(Err(err))
} else {
Poll::Ready(Ok(None))
}
} else if st.flags.is_read_eof() {
let notified = st.flags.take_read_notified();
if notified && st.flags.is_read_ready() {
Poll::Ready(Ok(Some(())))
} else {
Poll::Ready(Ok(None))
}
} else if st.flags.take_read_notified() {
Poll::Ready(Ok(Some(())))
} else {
st.flags.set_read_notify();
let _ = self.poll_read_more(cx);
st.dispatch_task.register(cx.waker());
Poll::Pending
}
}
#[inline]
pub fn poll_recv<U>(
&self,
codec: &U,
cx: &mut Context<'_>,
) -> Poll<Result<U::Item, RecvError<U>>>
where
U: Decoder,
{
let decoded = self.poll_recv_decode(codec, cx)?;
if let Some(item) = decoded.item {
Poll::Ready(Ok(item))
} else {
Poll::Pending
}
}
#[inline]
pub fn poll_recv_decode<U>(
&self,
codec: &U,
cx: &mut Context<'_>,
) -> Result<Decoded<U::Item>, RecvError<U>>
where
U: Decoder,
{
let st = self.st();
st.flags.unset_read_ready();
let closed = st.flags.is_stopping() || st.flags.is_terminating();
if !closed {
if st.flags.check_dispatcher_timeout() {
return Err(RecvError::Timeout);
} else if st.flags.is_wr_backpressure() {
return Err(RecvError::WriteBackpressure);
}
}
let decoded = self
.decode_item(codec)
.map_err(|err| RecvError::Decoder(err))?;
if decoded.item.is_some() {
Ok(decoded)
} else if st.flags.is_stopping() || st.flags.is_terminating() {
Err(RecvError::PeerGone(st.error()))
} else {
match self.poll_read_more(cx) {
Poll::Pending | Poll::Ready(Ok(Some(()))) => {
#[cfg(feature = "trace")]
if decoded.remains != 0 {
log::trace!("{}: Not enough data to decode next frame", self.tag());
}
Ok(decoded)
}
Poll::Ready(Err(e)) => Err(RecvError::PeerGone(Some(e))),
Poll::Ready(Ok(None)) => Err(RecvError::PeerGone(None)),
}
}
}
#[inline]
pub fn poll_flush(&self, cx: &mut Context<'_>, full: bool) -> Poll<io::Result<()>> {
let st = self.st();
st.buffer.process_write_buf_force(self)?;
self.consolidate_write_state(false)?;
let len = st.write_outstanding();
if len > 0 {
if st.flags.is_peer_gone() {
return Poll::Ready(Err(st.error_or_disconnected()));
} else if full {
st.flags.set_wants_write_flush();
st.dispatch_task.register(cx.waker());
return Poll::Pending;
} else if st.flags.is_wr_backpressure() {
if !st.should_disable_wr_backpressure(len) {
st.dispatch_task.register(cx.waker());
return Poll::Pending;
}
} else if st.is_wr_backpressure_needed(len) {
st.flags.set_wr_backpressure();
st.dispatch_task.register(cx.waker());
return Poll::Pending;
}
}
if st.flags.is_peer_gone() && !st.flags.is_write_flush() {
Poll::Ready(Err(st.error_or_disconnected()))
} else {
st.flags.unset_wr_backpressure_and_flush();
Poll::Ready(Ok(()))
}
}
#[inline]
pub fn poll_shutdown(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let st = self.st();
if st.flags.is_closed() {
if let Some(err) = st.error() {
Poll::Ready(Err(err))
} else {
Poll::Ready(Ok(()))
}
} else {
if !st.flags.is_terminating() && !st.flags.is_stopping_filters() {
st.start_shutdown();
}
if st.flags.is_read_paused() {
st.flags.unset_read_paused();
st.wake_read_task();
}
st.dispatch_task.register(cx.waker());
Poll::Pending
}
}
#[inline]
pub fn poll_read_pause(&self, cx: &mut Context<'_>) -> Poll<IoStatusUpdate> {
let st = self.st();
if !st.flags.is_read_paused() {
st.wake_read_task();
st.flags.set_read_paused();
}
self.poll_status_update(cx)
}
#[inline]
pub fn poll_status_update(&self, cx: &mut Context<'_>) -> Poll<IoStatusUpdate> {
let st = self.st();
st.dispatch_task.register(cx.waker());
if st.flags.is_peer_gone() {
Poll::Ready(IoStatusUpdate::PeerGone(st.error()))
} else if st.flags.check_dispatcher_timeout() {
Poll::Ready(IoStatusUpdate::Timeout)
} else if st.flags.is_wr_backpressure() {
if st.should_disable_wr_backpressure(st.write_outstanding()) {
st.flags.unset_wr_backpressure();
Poll::Pending
} else {
Poll::Ready(IoStatusUpdate::WriteBackpressure)
}
} else {
Poll::Pending
}
}
#[inline]
pub fn register_dispatch(&self, cx: &mut Context<'_>) {
self.st().dispatch_task.register(cx.waker());
}
}
impl<F> AsRef<IoRef> for Io<F> {
#[inline]
fn as_ref(&self) -> &IoRef {
self.io_ref()
}
}
impl<F> Eq for Io<F> {}
impl<F> PartialEq for Io<F> {
#[inline]
fn eq(&self, other: &Self) -> bool {
self.io_ref().eq(other.io_ref())
}
}
impl<F> hash::Hash for Io<F> {
#[inline]
fn hash<H: hash::Hasher>(&self, state: &mut H) {
self.io_ref().hash(state);
}
}
impl<F> fmt::Debug for Io<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Io").field("state", self.st()).finish()
}
}
impl<F> ops::Deref for Io<F> {
type Target = IoRef;
#[inline]
fn deref(&self) -> &Self::Target {
self.io_ref()
}
}
impl<F> Drop for Io<F> {
fn drop(&mut self) {
let st = self.st();
self.stop_timer();
let in_use = st.filter.is_set() && st.buffer.is_borrowed();
if st.filter.is_set() {
if in_use {
st.force_close_connection();
st.filter.leak();
} else {
if !st.flags.is_closed() {
log::trace!("{}: Io is dropped, terminate connection", st.tag());
}
if st.write_outstanding() == 0 {
st.terminate_connection(None);
} else {
st.force_close_connection();
}
st.filter.drop_filter::<F>();
}
st.buffer.release(self.io_ref().cfg());
drop(st.extensions.take_callbacks());
}
IoManager::unregister(self.io_ref());
assert!(
!in_use || std::thread::panicking(),
"{}: Io is dropped while its filter is in use",
st.tag()
);
}
}
fn as_u32(v: usize) -> u32 {
u32::try_from(v).unwrap_or(u32::MAX)
}
#[cfg(test)]
mod tests {
use std::{cell::Cell, rc::Rc};
use ntex_bytes::{BufMut, BytePages, Bytes, BytesMut};
use ntex_codec::BytesCodec;
use ntex_util::{future::lazy, time::Millis, time::sleep, time::timeout};
use super::*;
use crate::waiters::{TAG_DISCONNECT, WaiterEntry};
use crate::{
FilterBuf, IoContext, IoTaskStatus, Readiness, Waiter, ops::Iops, testing::IoTest,
};
use std::pin::Pin;
const BIN: &[u8] = b"GET /test HTTP/1\r\n\r\n";
const TEXT: &str = "GET /test HTTP/1\r\n\r\n";
const BIN2: &[u8] = b"12345678901234561234567890123456";
#[ntex::test]
async fn test_basics() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let server = Io::from(server);
assert!(server.eq(&server));
assert!(server.io_ref().eq(server.io_ref()));
}
#[ntex::test]
async fn test_recv() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let server = Io::new(server, SharedCfg::new("SRV"));
server.st().notify_timeout();
let err = server.recv(&BytesCodec).await.err().unwrap();
assert!(format!("{err:?}").contains("Timeout"));
client.write(TEXT);
server.st().flags.set_wr_backpressure();
let item = server.recv(&BytesCodec).await.ok().unwrap().unwrap();
assert_eq!(item, TEXT);
}
#[ntex::test]
async fn test_drop_releases_callbacks() {
struct Cb(#[allow(dead_code)] IoRef);
impl crate::IoCallbacks for Cb {
fn before_processing(&self, _: &IoRef) {}
fn after_processing(&self, _: &IoRef) {}
}
let (client, server) = IoTest::create();
let server = Io::new(server, SharedCfg::new("SRV"));
server.register_filter_callbacks(Cb(server.get_ref()));
let state = Rc::downgrade(&server.io_ref().0);
drop(server);
client.close().await;
sleep(Millis(50)).await;
assert!(state.upgrade().is_none());
}
#[ntex::test]
async fn test_callbacks_not_registered_after_drop_or_close() {
struct Cb(#[allow(dead_code)] IoRef, Rc<Cell<usize>>);
impl crate::IoCallbacks for Cb {
fn before_processing(&self, _: &IoRef) {
self.1.set(self.1.get() + 1);
}
fn after_processing(&self, _: &IoRef) {}
}
let (client, server) = IoTest::create();
let server = Io::new(server, SharedCfg::new("SRV"));
let io = server.get_ref();
let state = Rc::downgrade(&io.0);
drop(server);
io.register_filter_callbacks(Cb(io.clone(), Rc::default()));
drop(io);
client.close().await;
sleep(Millis(50)).await;
assert!(state.upgrade().is_none());
let (client, server) = IoTest::create();
let server = Io::new(server, SharedCfg::new("SRV"));
client.close().await;
server.close();
let _ = server.shutdown().await;
assert!(server.is_closed());
let calls = Rc::new(Cell::new(0));
server.register_filter_callbacks(Cb(server.get_ref(), calls.clone()));
server.with_callbacks(|cb| cb.before_processing(&server));
assert_eq!(calls.get(), 0);
}
#[ntex::test]
async fn test_stop_timer_clears_timeout_notification() {
let (_client, server) = IoTest::create();
let server = Io::new(server, SharedCfg::new("SRV"));
server.start_timer(ntex_util::time::Seconds(10));
server.notify_timeout();
server.stop_timer();
assert!(lazy(|cx| server.poll_status_update(cx)).await.is_pending());
}
#[ntex::test]
async fn test_read() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let server = Io::new(server, SharedCfg::new("SRV"));
client.write(b"1234");
let mut buf: [u8; 4] = [0, 0, 0, 0];
server.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"1234");
let fut = ntex_rt::spawn(async move {
let mut buf: [u8; 4] = [0, 0, 0, 0];
let err = server.read_exact(&mut buf).await.unwrap_err();
(server, err)
});
client.close().await;
let (server, err) = fut.await.unwrap();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
let err = server.read_exact(&mut [0]).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[ntex::test]
async fn test_read_partial_eof() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let server = Io::new(server, SharedCfg::new("SRV"));
client.write(b"12");
client.close().await;
let err = server.read_exact(&mut [0; 4]).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
let mut buf = [0; 2];
server.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"12");
let err = server.read_exact(&mut [0]).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[ntex::test]
async fn test_send() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let server = Io::from(server);
assert!(server.eq(&server));
server
.send(Bytes::from_static(BIN), &BytesCodec)
.await
.ok()
.unwrap();
let item = client.read_any();
assert_eq!(item, TEXT);
}
#[ntex::test]
async fn read() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_read_buf(8, 4)),
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
assert!(io.st().dispatch_task.is_set());
let ctx = IoContext::new(io.get_ref());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(io.st().read_task.is_set());
assert!(!io.st().flags.is_read_ready());
assert!(!io.st().flags.is_rd_backpressure());
assert!(!io.is_rd_backpressure());
assert!(!io.is_wr_backpressure());
ctx.release_read_buf(
BytesMut::copy_from_slice(b"1234567890"),
Poll::Ready(Ok(10)),
);
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_paused());
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_rd_backpressure());
assert!(io.is_rd_backpressure());
assert!(!io.is_wr_backpressure());
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert_eq!(io.with_read_dst(|buf| buf.split_to(1)), b"1");
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_rd_backpressure());
assert!(io.st().read_task.is_set());
assert_eq!(io.with_read_dst(|buf| buf.split_to(1)), b"2");
assert!(io.st().flags.is_rd_backpressure());
assert_eq!(io.with_read_dst(|buf| buf.split_to(1)), b"3");
assert!(io.st().flags.is_rd_backpressure());
assert!(io.st().flags.is_read_paused());
assert_eq!(io.with_read_dst(|buf| buf.split_to(3)), b"456");
assert!(!io.st().flags.is_read_paused());
assert!(!io.st().flags.is_read_ready());
assert!(!io.st().flags.is_rd_backpressure());
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
lazy(|cx| io.register_dispatch(cx)).await;
ctx.release_read_buf(BytesMut::copy_from_slice(b"1234"), Poll::Ready(Ok(4)));
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_paused());
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_rd_backpressure());
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert_eq!(io.with_read_dst(|buf| buf.split_to(4)), b"7890");
assert!(!io.st().flags.is_rd_backpressure());
lazy(|cx| io.register_dispatch(cx)).await;
ctx.release_read_buf(BytesMut::copy_from_slice(b"567"), Poll::Ready(Ok(3)));
assert!(!io.st().flags.is_read_paused());
assert!(io.st().flags.is_read_ready());
assert!(!io.st().flags.is_rd_backpressure());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert_eq!(io.with_read_dst(BytesMut::take), b"1234567");
assert!(!io.st().flags.is_read_paused());
assert!(!io.st().flags.is_read_ready());
assert!(io.st().read_task.is_set());
io.terminate();
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
}
#[ntex::test]
async fn only_force_close_reports_terminate() {
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ctx = IoContext::new(io.get_ref());
ctx.stop(Some(io::Error::other("transport failed")));
assert!(io.st().flags.is_terminating());
assert!(!io.st().flags.is_force_closing());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ctx = IoContext::new(io.get_ref());
io.close();
io.st().filters_stopped();
assert!(io.st().flags.is_stopping());
ctx.stop(Some(io::Error::other("transport failed")));
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ctx = IoContext::new(io.get_ref());
io.terminate();
assert!(io.st().flags.is_force_closing());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
}
#[ntex::test]
async fn drop_closes_gracefully_once_output_is_flushed() {
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ioref = io.get_ref();
let ctx = IoContext::new(io.get_ref());
assert_eq!(io.st().write_outstanding(), 0);
drop(io);
assert!(!ioref.0.flags.is_force_closing());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
}
#[ntex::test]
async fn drop_aborts_when_output_would_be_lost() {
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ioref = io.get_ref();
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"not delivered").unwrap();
assert_ne!(io.st().write_outstanding(), 0);
drop(io);
assert!(ioref.0.flags.is_force_closing());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
}
#[ntex::test]
async fn drop_releases_buffers() {
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ioref = io.get_ref();
let ctx = IoContext::new(io.get_ref());
ctx.release_read_buf(BytesMut::copy_from_slice(b"unread"), Poll::Ready(Ok(6)));
io.encode_slice(b"not delivered").unwrap();
assert_eq!(ioref.0.buffer.read_dst_size(), 6);
assert_ne!(ioref.0.buffer.write_buf_size(), 0);
drop(io);
assert_eq!(ioref.0.buffer.read_dst_size(), 0);
assert!(ioref.0.buffer.get_read_buf().is_none());
assert_eq!(ioref.0.buffer.write_buf_size(), 0);
ctx.release_read_buf(BytesMut::copy_from_slice(b"late"), Poll::Ready(Ok(4)));
assert!(ioref.0.buffer.get_read_buf().is_none());
ctx.with_read_buf(|buf| {
buf.extend_from_slice(b"late");
Poll::Ready(Ok(4))
});
assert!(ioref.0.buffer.get_read_buf().is_none());
assert_eq!(ioref.0.buffer.read_dst_size(), 0);
}
#[ntex::test]
async fn force_close_survives_filter_replacement() {
let io = Io::new(IoTest::create().0, SharedCfg::new("SRV"));
let ctx = IoContext::new(io.get_ref());
io.terminate();
drop(io);
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
}
#[ntex::test]
async fn read_notify() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_read_buf(8, 4)),
);
assert!(!io.st().flags.is_read_notify());
assert!(lazy(|cx| io.poll_read_notify(cx)).await.is_pending());
assert!(io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_notify());
let ctx = IoContext::new(io.get_ref());
ctx.release_read_buf(BytesMut::copy_from_slice(b"1"), Poll::Ready(Ok(1)));
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_read_notify());
assert!(io.st().flags.is_read_notified());
let res = lazy(|cx| io.poll_read_notify(cx)).await;
assert!(matches!(res, Poll::Ready(Ok(Some(())))));
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_ready());
assert!(lazy(|cx| io.poll_read_notify(cx)).await.is_pending());
assert!(io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_notify());
assert!(io.st().flags.is_read_ready());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
ctx.release_read_buf(BytesMut::copy_from_slice(b"2345678"), Poll::Ready(Ok(7)));
assert!(io.st().flags.is_rd_backpressure());
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_read_notify());
assert!(io.st().flags.is_read_notified());
let res = lazy(|cx| io.poll_read_notify(cx)).await;
assert!(matches!(res, Poll::Ready(Ok(Some(())))));
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert!(io.st().read_task.is_set());
assert!(lazy(|cx| io.poll_read_notify(cx)).await.is_pending());
assert!(!io.st().flags.is_rd_backpressure());
assert!(!io.st().flags.is_read_ready());
assert!(!io.st().flags.is_read_paused());
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
ctx.release_read_buf(BytesMut::copy_from_slice(b"1"), Poll::Ready(Ok(1)));
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().flags.is_read_ready());
assert!(io.st().flags.is_read_notify());
assert!(io.st().flags.is_read_paused());
assert!(io.st().flags.is_rd_backpressure());
assert!(io.st().flags.is_read_notified());
assert!(matches!(
lazy(|cx| io.poll_read_notify(cx)).await,
Poll::Ready(Ok(Some(())))
));
io.terminate();
let res = lazy(|cx| io.poll_read_notify(cx)).await;
assert!(matches!(res, Poll::Ready(Ok(None))), "{res:?}");
}
#[ntex::test]
async fn read_more() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::from(server);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
client.write(TEXT);
assert_eq!(io.read_more().await.unwrap(), Some(()));
assert!(matches!(
lazy(|cx| io.poll_read_more(cx)).await,
Poll::Ready(Ok(Some(())))
));
let item = io.with_read_dst(BytesMut::take);
assert_eq!(item, Bytes::from_static(BIN));
client.write(TEXT);
sleep(Millis(50)).await;
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_ready());
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_ready());
}
#[ntex::test]
async fn read_backpressure() {
let (client, server) = IoTest::create();
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_read_buf(64, 32)),
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
client.write(BIN2);
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_read_ready());
assert!(io.flags().is_rd_backpressure());
let _item = io.recv(&BytesCodec).await.ok().unwrap().unwrap();
assert!(!io.flags().is_read_ready());
assert!(!io.flags().is_rd_backpressure());
client.write(BIN2);
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_read_ready());
assert!(io.flags().is_rd_backpressure());
assert_eq!(io.read_more().await.unwrap(), Some(()));
}
#[ntex::test]
async fn read_src_releases_read_backpressure() {
let (client, server) = IoTest::create();
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_read_buf(64, 32)),
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
client.write(BIN2);
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_rd_backpressure());
let len = io.get_ref().with_read_src(|buf| {
let len = buf.len();
buf.clear();
len
});
assert!(len > 0);
assert!(!io.flags().is_rd_backpressure());
assert!(!io.flags().is_read_paused());
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_read_ready());
}
#[ntex::test]
async fn with_buf_releases_read_backpressure() {
let (client, server) = IoTest::create();
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_read_buf(64, 32)),
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
client.write(BIN2);
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_rd_backpressure());
let len = io
.get_ref()
.with_buf(|buf| {
buf.with_read_buffers(|_, dst| {
let len = dst.len();
dst.clear();
len
})
})
.unwrap();
assert!(len > 0);
assert!(!io.flags().is_rd_backpressure());
assert!(!io.flags().is_read_paused());
client.write(BIN2);
sleep(Millis(50)).await;
assert!(io.flags().is_read_ready());
}
#[ntex::test]
async fn write() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
assert!(lazy(|cx| io.poll_status_update(cx)).await.is_pending());
assert!(io.st().dispatch_task.is_set());
assert!(io.st().flags.is_direct_wr_enabled());
let ctx = IoContext::new(io.get_ref());
assert_eq!(lazy(|cx| ctx.poll_write_ready(cx)).await, Poll::Pending);
assert!(io.st().write_task.is_set());
assert!(io.st().flags.is_write_paused());
assert!(!io.st().flags.is_wr_backpressure());
io.with_write_src(|buf| buf.put_slice(b"1234")).unwrap();
assert_eq!(lazy(|cx| ctx.poll_write_ready(cx)).await, Poll::Pending);
assert!(io.st().flags.is_write_paused());
assert!(io.st().flags.is_wr_send_scheduled());
assert!(!io.st().flags.is_wr_backpressure());
assert!(io.st().dispatch_task.is_set());
io.with_write_src(|buf| buf.put_slice(b"5678")).unwrap();
assert!(io.st().flags.is_wr_backpressure());
assert!(!io.st().dispatch_task.is_set());
assert!(io.st().write_task.is_set());
assert!(matches!(
lazy(|cx| io.poll_status_update(cx)).await,
Poll::Ready(IoStatusUpdate::WriteBackpressure)
));
assert!(lazy(|cx| io.poll_flush(cx, false)).await.is_pending());
assert!(!io.st().flags.is_write_flush());
Iops::run();
assert!(!io.st().flags.is_wr_send_scheduled());
assert!(!io.st().flags.is_write_paused());
assert!(!io.st().write_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(4).freeze()), b"1234");
assert_eq!(ctx.update_write_status(Ok(4)), IoTaskStatus::Io);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.st().flags.is_write_paused());
assert!(io.st().flags.is_wr_backpressure());
assert!(lazy(|cx| io.poll_status_update(cx)).await.is_pending());
assert!(!io.st().flags.is_wr_backpressure());
assert!(lazy(|cx| io.poll_status_update(cx)).await.is_pending());
assert!(matches!(
lazy(|cx| io.poll_flush(cx, false)).await,
Poll::Ready(Ok(()))
));
io.with_write_src(|buf| buf.put_slice(b"1234")).unwrap();
assert!(lazy(|cx| io.poll_flush(cx, true)).await.is_pending());
assert!(io.st().flags.is_write_flush());
assert!(io.st().flags.is_wr_backpressure());
Iops::run();
assert_eq!(ctx.with_write_dst(BytePages::freeze), b"56781234");
assert!(!io.st().flags.is_wr_send_scheduled());
assert_eq!(ctx.update_write_status(Ok(8)), IoTaskStatus::Pause);
assert!(io.st().flags.is_write_paused());
assert!(io.st().flags.is_write_flush());
assert!(io.st().flags.is_wr_backpressure());
assert!(!io.st().dispatch_task.is_set());
assert!(matches!(
lazy(|cx| io.poll_flush(cx, false)).await,
Poll::Ready(Ok(()))
));
assert!(!io.st().flags.is_write_flush());
assert!(!io.st().flags.is_wr_backpressure());
io.terminate();
assert!(!io.st().write_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Terminate)
);
let Poll::Ready(Err(err)) = lazy(|cx| io.poll_flush(cx, false)).await else {
panic!()
};
assert_eq!(err.kind(), io::ErrorKind::NotConnected);
assert!(matches!(
lazy(|cx| io.poll_status_update(cx)).await,
Poll::Ready(IoStatusUpdate::PeerGone(None))
));
}
#[ntex::test]
async fn local_shutdown_reports_peer_gone_without_error() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::from(server);
io.shutdown().await.unwrap();
assert!(!io.is_active());
assert!(matches!(
lazy(|cx| io.poll_status_update(cx)).await,
Poll::Ready(IoStatusUpdate::PeerGone(None))
));
}
#[ntex::test]
async fn set_config_updates_eager_write_support() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(0)),
);
assert!(!io.st().flags.is_direct_wr_enabled());
unsafe {
io.set_config(SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(1024)));
}
assert!(io.st().flags.is_direct_wr_enabled());
assert_eq!(io.cfg().write_buf_threshold(), 1024);
unsafe {
io.set_config(SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(0)));
}
assert!(!io.st().flags.is_direct_wr_enabled());
assert_eq!(io.cfg().write_buf_threshold(), 0);
}
#[ntex::test]
async fn eager_write_uses_updated_buffer_size() {
#[derive(Debug)]
struct DirectWrite;
impl IoStream for DirectWrite {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DirectWrite {
fn write(&self, ctx: &IoContext) {
let n = ctx.with_write_dst(|buf| {
let n = buf.len();
buf.clear();
n
});
let _ = ctx.update_write_status(Ok(n));
}
}
let io = Io::new(
DirectWrite,
SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(1).set_write_buf(8)),
);
io.encode_slice(BIN2).unwrap();
assert_eq!(io.st().buffer.write_buf_size(), 0);
assert!(io.flags().is_write_paused());
assert!(!io.flags().is_wr_backpressure());
assert!(!io.st().flags.is_wr_send_scheduled());
}
#[ntex::test]
async fn eager_write_reports_transport_error() {
#[derive(Debug)]
struct FailedWrite;
impl IoStream for FailedWrite {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for FailedWrite {
fn write(&self, ctx: &IoContext) {
ctx.update_write_status(Err(io::Error::new(
io::ErrorKind::ConnectionReset,
"connection reset",
)));
}
}
let io = Io::new(
FailedWrite,
SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(1)),
);
let err = io.encode_slice(BIN2).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
assert!(io.st().flags.is_terminating());
}
#[ntex::test]
async fn terminate_during_eager_write_releases_transport() {
#[derive(Debug)]
struct FailedWrite(Rc<Cell<bool>>);
impl Drop for FailedWrite {
fn drop(&mut self) {
self.0.set(true);
}
}
impl IoStream for FailedWrite {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for FailedWrite {
fn write(&self, ctx: &IoContext) {
ctx.update_write_status(Err(io::Error::new(
io::ErrorKind::ConnectionReset,
"connection reset",
)));
}
}
let dropped = Rc::new(Cell::new(false));
let io = Io::new(
FailedWrite(dropped.clone()),
SharedCfg::new("SRV").add(IoConfig::new().set_write_buf_threshold(1)),
);
io.encode_slice(BIN2).unwrap_err();
assert!(io.st().flags.is_terminating());
assert!(dropped.get());
assert!(io.st().handle.take().is_none());
}
#[ntex::test]
async fn write_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(16)),
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
assert!(io.flags().is_write_paused());
assert!(!io.flags().is_wr_backpressure());
assert!(!io.is_wr_backpressure());
io.encode_slice(BIN2).unwrap();
assert!(Iops::is_registered(&io));
assert!(io.flags().is_wr_backpressure());
client.remote_buffer_cap(1024);
let item = client.read().await.unwrap();
assert_eq!(item, BIN2);
assert!(io.flags().is_wr_backpressure());
assert!(lazy(|cx| io.poll_status_update(cx)).await.is_pending());
assert!(!io.flags().is_wr_backpressure());
assert!(matches!(
lazy(|cx| io.poll_flush(cx, false)).await,
Poll::Ready(Ok(()))
));
assert!(!io.flags().is_wr_backpressure());
}
#[ntex::test]
async fn partial_flush_keeps_write_backpressure_until_half_watermark() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"12345678").unwrap();
assert!(io.flags().is_wr_backpressure());
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(1).len()), 1);
assert_eq!(ctx.update_write_status(Ok(1)), IoTaskStatus::Io);
assert!(lazy(|cx| io.poll_flush(cx, false)).await.is_pending());
assert!(io.flags().is_wr_backpressure());
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(3).len()), 3);
assert_eq!(ctx.update_write_status(Ok(3)), IoTaskStatus::Io);
assert!(matches!(
lazy(|cx| io.poll_flush(cx, false)).await,
Poll::Ready(Ok(()))
));
assert!(!io.flags().is_wr_backpressure());
}
#[ntex::test]
async fn write_ready_waits_for_release_threshold() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let ctx = IoContext::new(io.get_ref());
assert!(io.write_ready().await.is_ok());
io.encode_slice(b"12345678").unwrap();
assert!(io.flags().is_wr_backpressure());
let done = Rc::new(Cell::new(0));
for _ in 0..2 {
let (io, done) = (io.get_ref(), done.clone());
ntex_util::spawn(async move {
io.write_ready().await.unwrap();
done.set(done.get() + 1);
});
}
sleep(Millis(10)).await;
assert_eq!(done.get(), 0);
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(1).len()), 1);
assert_eq!(ctx.update_write_status(Ok(1)), IoTaskStatus::Io);
sleep(Millis(10)).await;
assert_eq!(done.get(), 0);
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(3).len()), 3);
assert_eq!(ctx.update_write_status(Ok(3)), IoTaskStatus::Io);
sleep(Millis(10)).await;
assert_eq!(done.get(), 2);
assert!(io.flags().is_wr_backpressure());
assert!(io.write_ready().await.is_ok());
}
#[ntex::test]
async fn write_ready_released_during_full_flush() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"12345678").unwrap();
assert!(lazy(|cx| io.poll_flush(cx, true)).await.is_pending());
assert!(io.flags().is_write_flush());
let done = Rc::new(Cell::new(false));
let (io2, done2) = (io.get_ref(), done.clone());
ntex_util::spawn(async move {
io2.write_ready().await.unwrap();
done2.set(true);
});
sleep(Millis(10)).await;
assert!(!done.get());
assert_eq!(ctx.with_write_dst(|buf| buf.split_to(4).len()), 4);
assert_eq!(ctx.update_write_status(Ok(4)), IoTaskStatus::Io);
sleep(Millis(10)).await;
assert!(done.get());
}
#[ntex::test]
async fn write_ready_fails_on_disconnect() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
io.encode_slice(b"12345678").unwrap();
assert!(io.flags().is_wr_backpressure());
let res = Rc::new(Cell::new(None));
let (io2, res2) = (io.get_ref(), res.clone());
ntex_util::spawn(async move {
res2.set(Some(io2.write_ready().await.is_err()));
});
sleep(Millis(10)).await;
assert_eq!(res.get(), None);
io.terminate();
sleep(Millis(10)).await;
assert_eq!(res.get(), Some(true));
assert!(io.write_ready().await.is_err());
}
#[ntex::test]
async fn waiter() {
let (client, server) = IoTest::create();
let io = Io::from(server);
let (mut s1, mut s2, mut s3) = (Waiter::new(&io, 7), Waiter::new(&io, 7), io.waiter(8));
assert!(lazy(|cx| Pin::new(&mut s1).poll(cx)).await.is_pending());
assert!(lazy(|cx| Pin::new(&mut s2).poll(cx)).await.is_pending());
assert!(lazy(|cx| Pin::new(&mut s3).poll(cx)).await.is_pending());
assert!(lazy(|cx| Pin::new(&mut s1).poll(cx)).await.is_pending());
assert_eq!(io.get_ref().0.extensions.wakers_len(), 3);
io.wake(7);
assert!(lazy(|cx| Pin::new(&mut s1).poll(cx)).await.is_ready());
assert!(lazy(|cx| Pin::new(&mut s2).poll(cx)).await.is_ready());
assert!(lazy(|cx| Pin::new(&mut s3).poll(cx)).await.is_pending());
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
drop(s3);
assert_eq!(io.get_ref().0.extensions.wakers_len(), 0);
assert!(lazy(|cx| Pin::new(&mut s1).poll(cx)).await.is_pending());
let mut s1 = s1.into_static();
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
let s4 = s1.clone();
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
drop(s4);
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
let res = Rc::new(Cell::new(false));
let res2 = res.clone();
ntex_util::spawn(async move {
(&mut s1).await;
res2.set(true);
});
sleep(Millis(10)).await;
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
client.close().await;
timeout(Millis(1000), io.shutdown())
.await
.expect("stream shutdown did not complete")
.unwrap();
sleep(Millis(10)).await;
assert!(res.get(), "waiter was not woken on disconnect");
(&mut s2).await;
drop(s2);
assert_eq!(io.get_ref().0.extensions.wakers_len(), 0);
}
#[ntex::test]
async fn waiter_poll_ready() {
let (_client, server) = IoTest::create();
let io = Io::from(server);
let waiter = io.waiter(3);
let ext = &io.get_ref().0.extensions;
io.wake(3);
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
assert_eq!(ext.wakers_len(), 1);
io.wake(3);
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_ready());
assert_eq!(ext.wakers_len(), 0);
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
assert_eq!(ext.wakers_len(), 1);
io.wake(3);
waiter.await;
assert_eq!(ext.wakers_len(), 0);
}
#[ntex::test]
async fn wake_reserved_tag_is_ignored() {
let (_client, server) = IoTest::create();
let io = Io::from(server);
let mut waiter = io.on_disconnect();
assert!(lazy(|cx| Pin::new(&mut waiter).poll(cx)).await.is_pending());
io.wake(TAG_DISCONNECT);
io.wake(TAG_WRITE);
assert!(lazy(|cx| Pin::new(&mut waiter).poll(cx)).await.is_pending());
assert_eq!(io.get_ref().0.extensions.wakers_len(), 1);
}
#[cfg(debug_assertions)]
#[ntex::test]
#[should_panic(expected = "reserved")]
async fn waiter_reserved_tag() {
let (_client, server) = IoTest::create();
let io = Io::from(server);
let _waiter = Waiter::new(&io, TAG_DISCONNECT);
}
#[cfg(debug_assertions)]
#[ntex::test]
#[should_panic(expected = "reserved")]
async fn waiter_reserved_tag_ioref() {
let (_client, server) = IoTest::create();
let io = Io::from(server);
let _waiter = io.waiter(TAG_WRITE);
}
#[ntex::test]
async fn dropped_waiters_are_removed() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let ext = &io.get_ref().0.extensions;
for _ in 0..4 {
let waiter = io.on_disconnect();
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
assert_eq!(ext.wakers_len(), 1);
}
assert_eq!(ext.wakers_len(), 0);
io.encode_slice(b"12345678").unwrap();
assert!(io.flags().is_wr_backpressure());
for _ in 0..4 {
let mut fut = std::pin::pin!(io.write_ready());
assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
assert_eq!(ext.wakers_len(), 1);
}
assert_eq!(ext.wakers_len(), 0);
let waiter = io.on_disconnect();
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
let mut fut = std::pin::pin!(io.write_ready());
assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_pending());
assert_eq!(ext.wakers_len(), 2);
io.terminate();
assert_eq!(ext.wakers_len(), 1);
assert!(lazy(|cx| fut.as_mut().poll(cx)).await.is_ready());
sleep(Millis(50)).await;
assert_eq!(ext.wakers_len(), 0);
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_ready());
}
#[test]
fn woken_waiter_slot_is_reused() {
let ext = Extensions::default();
let waker = std::task::Waker::noop();
let (a, b) = (WaiterEntry::new(TAG_WRITE), WaiterEntry::new(TAG_WRITE));
ext.register_waker(&a, waker);
ext.wake(TAG_WRITE);
ext.register_waker(&b, waker);
ext.register_waker(&a, waker);
assert_eq!(ext.wakers_len(), 2);
ext.remove_waker(&a);
assert_eq!(ext.wakers_len(), 1);
ext.wake(TAG_WRITE);
ext.register_waker(&b, waker);
ext.register_waker(&a, waker);
ext.wake(TAG_WRITE);
ext.register_waker(&b, waker);
ext.remove_waker(&a);
assert_eq!(ext.wakers_len(), 1);
ext.remove_waker(&b);
assert_eq!(ext.wakers_len(), 0);
}
#[ntex::test]
async fn full_flush_waits_for_inflight_write() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default()),
);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"12345678").unwrap();
let page = ctx.with_write_dst(BytePages::take).unwrap();
assert_eq!(page.len(), 8);
assert_eq!(io.st().buffer.write_buf_size(), 0);
assert_eq!(io.st().write_outstanding(), 8);
assert!(lazy(|cx| io.poll_flush(cx, true)).await.is_pending());
assert_eq!(ctx.update_write_status(Ok(page.len())), IoTaskStatus::Pause);
assert_eq!(io.st().write_outstanding(), 0);
assert!(matches!(
lazy(|cx| io.poll_flush(cx, true)).await,
Poll::Ready(Ok(()))
));
}
#[ntex::test]
async fn write_backpressure_counts_inflight_output() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"12345678").unwrap();
assert!(io.flags().is_wr_backpressure());
let page = ctx.with_write_dst(BytePages::take).unwrap();
assert_eq!(io.st().buffer.write_buf_size(), 0);
assert!(lazy(|cx| io.poll_flush(cx, false)).await.is_pending());
assert!(io.flags().is_wr_backpressure());
assert_eq!(ctx.update_write_status(Ok(page.len())), IoTaskStatus::Pause);
assert!(matches!(
lazy(|cx| io.poll_flush(cx, false)).await,
Poll::Ready(Ok(()))
));
assert!(!io.flags().is_wr_backpressure());
}
#[ntex::test]
async fn transport_shutdown_waits_for_inflight_write() {
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default()),
);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"12345678").unwrap();
let page = ctx.with_write_dst(BytePages::take).unwrap();
io.st().flags.enter_filters_stopping();
io.st().filters_stopped();
assert!(io.st().flags.is_stopping());
assert_eq!(lazy(|cx| ctx.poll_write_ready(cx)).await, Poll::Pending);
assert_eq!(ctx.update_write_status(Ok(page.len())), IoTaskStatus::Pause);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
);
}
#[ntex::test]
async fn shutdown_flushes_write_buf_with_read_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(8, 4)
.set_shutdown_timeout(ntex_util::time::Seconds(2)),
),
);
io.encode_slice(b"response-tail").unwrap();
sleep(Millis(50)).await;
assert_eq!(io.st().buffer.write_buf_size(), 13);
client.write("0123456789");
sleep(Millis(50)).await;
assert!(io.flags().is_read_paused());
assert!(io.flags().is_rd_backpressure());
assert_eq!(io.st().buffer.write_buf_size(), 13);
io.close();
sleep(Millis(50)).await;
client.remote_buffer_cap(1024);
let data = ntex_util::time::timeout(Millis(2000), client.read())
.await
.expect("write buffer was dropped during shutdown")
.unwrap();
assert_eq!(&data[..], b"response-tail");
ntex_util::time::timeout(Millis(4000), io.on_disconnect())
.await
.expect("io stream did not disconnect after flush");
}
#[ntex::test]
async fn peer_eof_allows_response_before_shutdown() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::from(server);
client.write("request");
client.close().await;
assert_eq!(
timeout(Millis(1000), io.recv(&BytesCodec))
.await
.expect("request was not decoded")
.unwrap(),
Some(Bytes::from_static(b"request"))
);
assert!(
timeout(Millis(1000), io.recv(&BytesCodec))
.await
.expect("EOF was not reported")
.unwrap()
.is_none()
);
assert!(!io.st().flags.is_closed());
io.encode(Bytes::from_static(b"response"), &BytesCodec)
.unwrap();
timeout(Millis(1000), io.shutdown())
.await
.expect("shutdown did not complete")
.unwrap();
assert_eq!(
timeout(Millis(1000), client.read())
.await
.expect("response was not flushed")
.unwrap(),
b"response"[..]
);
}
#[ntex::test]
async fn shutdown_waits_for_transport_stop() {
#[derive(Debug)]
struct DormantTransport;
impl IoStream for DormantTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DormantTransport {}
let io = Io::from(DormantTransport);
let ctx = IoContext::new(io.get_ref());
let waiter = io.on_disconnect();
io.st().flags.enter_transport_shutdown();
assert!(lazy(|cx| io.poll_shutdown(cx)).await.is_pending());
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
ctx.stopped(None);
assert!(matches!(
lazy(|cx| io.poll_shutdown(cx)).await,
Poll::Ready(Ok(()))
));
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_ready());
}
#[ntex::test]
async fn a_stopped_connection_keeps_its_buffered_input() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
client.write(TEXT);
let io = Io::from(server);
io.read_more().await.unwrap().unwrap();
assert_eq!(io.with_read_dst(|buf| buf.len()), BIN.len());
IoContext::new(io.get_ref()).stopped(None);
assert!(io.is_closed() && !io.is_active());
assert!(io.st().flags.is_stopping() && !io.st().flags.is_shutting_down_filters());
assert_eq!(
io.with_read_dst(|buf| buf.len()),
BIN.len(),
"input received before the transport went away was discarded"
);
}
#[ntex::test]
async fn transport_shutdown_drain_wakes_write_task() {
#[derive(Debug)]
struct DormantTransport;
impl IoStream for DormantTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DormantTransport {}
let io = Io::from(DormantTransport);
let ctx = IoContext::new(io.get_ref());
io.encode_slice(b"tail").unwrap();
io.st().flags.enter_filters_stopping();
io.st().flags.enter_transport_shutdown();
assert_eq!(io.st().buffer.write_buf_size(), 4);
assert!(matches!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Ready)
));
assert!(io.st().write_task.is_set());
let res = ctx.with_write_dst(|buf| {
let mut written = 0;
while let Some(page) = buf.take() {
written += page.len();
}
Ok(written)
});
assert_eq!(ctx.update_write_status(res), IoTaskStatus::Pause);
assert_eq!(io.st().write_outstanding(), 0);
assert!(!io.st().write_task.is_set());
assert!(matches!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Close)
));
}
#[ntex::test]
async fn termination_waits_for_transport_stop() {
#[derive(Debug)]
struct DormantTransport;
impl IoStream for DormantTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DormantTransport {}
let io = Io::from(DormantTransport);
let ctx = IoContext::new(io.get_ref());
let waiter = io.on_disconnect();
ctx.stop(Some(io::Error::new(
io::ErrorKind::ConnectionReset,
"connection reset",
)));
assert!(io.st().flags.is_terminating());
assert!(!io.st().flags.is_closed());
assert!(lazy(|cx| io.poll_shutdown(cx)).await.is_pending());
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_pending());
ctx.stopped(None);
let Poll::Ready(Err(err)) = lazy(|cx| io.poll_shutdown(cx)).await else {
panic!("shutdown did not report termination error");
};
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert!(lazy(|cx| waiter.poll_ready(cx)).await.is_ready());
}
#[ntex::test]
async fn send_buf_retains_termination_error() {
#[derive(Debug)]
struct DormantTransport;
impl IoStream for DormantTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DormantTransport {}
let io = Io::from(DormantTransport);
let ctx = IoContext::new(io.get_ref());
ctx.stop(Some(io::Error::new(
io::ErrorKind::ConnectionReset,
"connection reset",
)));
let err = io.get_ref().send_buf().unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
ctx.stopped(None);
let err = io.shutdown().await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
}
#[ntex::test]
async fn read_readiness_reports_termination_error() {
#[derive(Debug)]
struct DormantTransport;
impl IoStream for DormantTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for DormantTransport {}
let io = Io::from(DormantTransport);
let ctx = IoContext::new(io.get_ref());
ctx.stop(Some(io::Error::new(
io::ErrorKind::ConnectionReset,
"connection reset",
)));
let Poll::Ready(Err(err)) = lazy(|cx| io.poll_read_more(cx)).await else {
panic!("read request did not report termination error");
};
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
let Poll::Ready(Err(err)) = lazy(|cx| io.poll_read_notify(cx)).await else {
panic!("read notification did not report termination error");
};
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
let err = io.read_exact(&mut [0]).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::ConnectionReset);
assert_eq!(err.to_string(), "connection reset");
}
#[ntex::test]
async fn filter_shutdown_is_blocked_by_read_backpressure() {
#[derive(Debug)]
struct PendingShutdown;
impl FilterLayer for PendingShutdown {
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<()>> {
Ok(Poll::Pending)
}
}
let (_client, server) = IoTest::create();
let io = Io::new(server, SharedCfg::new("SRV")).add_filter(PendingShutdown);
io.st().flags.set_read_ready_and_backpressure();
io.st().flags.unset_read_ready();
assert!(lazy(|cx| io.poll_shutdown(cx)).await.is_pending());
let ctx = IoContext::new(io.get_ref());
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert!(io.st().flags.is_stopping());
}
#[ntex::test]
async fn intermediate_filter_output_reaches_transport() {
#[derive(Debug)]
struct Emit;
impl FilterLayer for Emit {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_read_buffers(|src, dst| {
if let Some(src) = src {
dst.extend_from_slice(src);
src.clear();
}
});
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"pong"));
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Ready(()))
}
}
#[derive(Debug)]
struct Passthrough;
impl FilterLayer for Passthrough {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_read_buffers(|src, dst| {
if let Some(src) = src {
dst.extend_from_slice(src);
src.clear();
}
});
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Ready(()))
}
}
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::from(server).add_filter(Passthrough).add_filter(Emit);
client.write("ping");
let _ = io.recv(&BytesCodec).await.unwrap();
sleep(Millis(50)).await;
assert!(client.read_any().starts_with(b"pong"));
}
#[ntex::test]
async fn read_pauses_while_read_output_is_not_drained() {
#[derive(Debug)]
struct Reply;
impl FilterLayer for Reply {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
let data = buf.with_read_buffers(|src, _| src.as_mut().map(BytesMut::take));
if let Some(data) = data {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(&data));
}
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Ready(()))
}
}
#[derive(Debug)]
struct Manual;
impl IoStream for Manual {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for Manual {}
let io = Io::new(
Manual,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
)
.add_filter(Reply);
let ctx = IoContext::new(io.get_ref());
let read = |data: &'static [u8]| {
ctx.with_read_buf(|buf| {
buf.extend_from_slice(data);
Poll::Ready(Ok(data.len()))
})
};
assert_eq!(read(b"ping"), IoTaskStatus::Io);
assert!(!io.st().flags.is_read_wr_backpressure());
assert_eq!(read(b"ping"), IoTaskStatus::Pause);
assert!(io.st().flags.is_read_wr_backpressure());
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert!(io.st().read_task.is_set());
let _ = ctx.with_write_dst(|buf| buf.split_to(8));
assert_eq!(io.st().write_outstanding(), 8);
let _ = ctx.update_write_status(Ok(0));
assert!(io.st().flags.is_read_wr_backpressure());
let _ = ctx.update_write_status(Ok(3));
assert!(io.st().flags.is_read_wr_backpressure());
assert!(io.st().read_task.is_set());
let _ = ctx.update_write_status(Ok(1));
assert!(!io.st().flags.is_read_wr_backpressure());
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
io.encode_slice(b"12345678").unwrap();
assert!(io.st().write_outstanding() >= 8);
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.st().flags.is_read_wr_backpressure());
}
#[ntex::test]
async fn read_output_pause_ignores_held_back_output() {
#[derive(Debug, Default)]
struct Reneg(Cell<bool>);
impl FilterLayer for Reneg {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
let data = buf.with_read_buffers(|src, _| src.as_mut().map(BytesMut::take));
match data.as_deref() {
Some(b"hello") => {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"handshake"));
}
Some(b"done") => {
self.0.set(true);
buf.with_write_buffers(BytePages::move_to);
}
_ => (),
}
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
if self.0.get() {
buf.with_write_buffers(BytePages::move_to);
}
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Ready(()))
}
}
#[derive(Debug)]
struct Manual;
impl IoStream for Manual {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for Manual {}
let io = Io::new(
Manual,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
)
.add_filter(Reneg::default());
let ctx = IoContext::new(io.get_ref());
let read = |data: &'static [u8]| {
ctx.with_read_buf(|buf| {
buf.extend_from_slice(data);
Poll::Ready(Ok(data.len()))
})
};
io.encode_slice(b"12345678").unwrap();
assert_eq!(io.st().write_outstanding(), 8);
assert_eq!(read(b"hello"), IoTaskStatus::Pause);
assert!(io.st().flags.is_read_wr_backpressure());
let _ = ctx.with_write_dst(|buf| buf.split_to(9));
let _ = ctx.update_write_status(Ok(9));
assert_eq!(io.st().write_outstanding(), 8);
assert!(!io.st().flags.is_read_wr_backpressure());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
let _ = read(b"done");
assert_eq!(ctx.with_write_dst(|buf| buf.len()), 8);
}
#[ntex::test]
async fn read_more_lifts_read_output_pause() {
#[derive(Debug)]
struct Reply;
impl FilterLayer for Reply {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
let data = buf.with_read_buffers(|src, _| src.as_mut().map(BytesMut::take));
if let Some(data) = data {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(&data));
}
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Ready(()))
}
}
#[derive(Debug)]
struct Manual;
impl IoStream for Manual {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for Manual {}
let io = Io::new(
Manual,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
)
.add_filter(Reply);
let ctx = IoContext::new(io.get_ref());
let read = |data: &'static [u8]| {
ctx.with_read_buf(|buf| {
buf.extend_from_slice(data);
Poll::Ready(Ok(data.len()))
})
};
assert_eq!(read(b"pingping"), IoTaskStatus::Pause);
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
assert!(!io.st().flags.is_read_wr_backpressure());
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert_eq!(read(b"ping"), IoTaskStatus::Pause);
assert_eq!(lazy(|cx| ctx.poll_read_ready(cx)).await, Poll::Pending);
assert!(lazy(|cx| io.poll_read_notify(cx)).await.is_pending());
assert!(!io.st().flags.is_read_wr_backpressure());
assert!(!io.st().read_task.is_set());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
}
#[ntex::test]
async fn read_pause_on_read_output_does_not_block_filter_shutdown() {
#[derive(Debug)]
struct Reply;
impl FilterLayer for Reply {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
let data = buf.with_read_buffers(|src, _| src.as_mut().map(BytesMut::take));
if let Some(data) = data {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(&data));
}
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
Ok(Poll::Pending)
}
}
#[derive(Debug)]
struct Manual;
impl IoStream for Manual {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for Manual {}
let io = Io::new(
Manual,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
)
.add_filter(Reply);
let ctx = IoContext::new(io.get_ref());
let read = |data: &'static [u8]| {
ctx.with_read_buf(|buf| {
buf.extend_from_slice(data);
Poll::Ready(Ok(data.len()))
})
};
assert_eq!(read(b"pingping"), IoTaskStatus::Pause);
assert!(io.st().flags.is_read_wr_backpressure());
assert!(lazy(|cx| io.poll_shutdown(cx)).await.is_pending());
assert!(io.st().flags.is_stopping_filters());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert_eq!(read(b"ping"), IoTaskStatus::Io);
}
#[ntex::test]
async fn peer_eof_completes_filter_shutdown() {
#[derive(Debug)]
struct PendingShutdown;
impl FilterLayer for PendingShutdown {
fn process_read_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
buf.with_write_buffers(BytePages::move_to);
Ok(Poll::Pending)
}
}
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(
server,
SharedCfg::new("SRV")
.add(IoConfig::default().set_shutdown_timeout(ntex_util::time::Seconds(30))),
)
.add_filter(PendingShutdown);
io.encode_slice(b"bye").unwrap();
let peer = client.clone();
drop(client);
assert!(io.read_more().await.unwrap().is_none());
assert!(io.st().flags.is_read_eof());
timeout(Millis(1000), io.shutdown())
.await
.expect("transport shutdown did not complete")
.unwrap();
assert!(io.st().flags.is_closed());
assert_eq!(peer.read_any(), Bytes::from_static(b"bye"));
}
#[derive(Debug)]
struct StuckShutdown(Cell<bool>);
impl FilterLayer for StuckShutdown {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_read_buffers(|src, dst| {
if let Some(src) = src {
dst.extend_from_slice(src);
src.clear();
}
});
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, _: &FilterBuf<'_>) -> io::Result<Poll<()>> {
self.0.set(true);
Ok(Poll::Pending)
}
}
#[ntex::test]
async fn transport_shutdown_pauses_read_task() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("before");
sleep(Millis(25)).await;
assert_eq!(
io.recv(&BytesCodec).await.unwrap().unwrap(),
b"before".as_ref()
);
client.remote_buffer_cap(0);
io.get_ref().with_write_dst(|b| b.extend_from_slice(b"out"));
io.st().flags.enter_filters_stopping();
io.st().flags.enter_transport_shutdown();
io.st().wake_read_task();
client.write("after");
sleep(Millis(50)).await;
assert!(!io.st().flags.is_closed());
assert_eq!(client.remote_buffer(|buf| buf.len()), 5);
}
#[derive(Default)]
struct WakeCounter(std::sync::atomic::AtomicUsize);
impl std::task::Wake for WakeCounter {
fn wake(self: std::sync::Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &std::sync::Arc<Self>) {
self.0.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
}
impl WakeCounter {
fn count(&self) -> usize {
self.0.load(std::sync::atomic::Ordering::Relaxed)
}
}
#[ntex::test]
async fn repeated_shutdown_polls_do_not_wake_tasks() {
use std::{sync::Arc, task::Waker};
let (_client, server) = IoTest::create();
let io = Io::new(server, SharedCfg::new("SRV")).add_filter(StuckShutdown(Cell::new(false)));
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
assert!(io.poll_shutdown(&mut cx).is_pending());
assert!(io.st().flags.is_stopping_filters());
let rd = Arc::new(WakeCounter::default());
let wr = Arc::new(WakeCounter::default());
io.st().read_task.register(&Waker::from(rd.clone()));
io.st().write_task.register(&Waker::from(wr.clone()));
for _ in 0..3 {
assert!(io.poll_shutdown(&mut cx).is_pending());
}
assert_eq!(rd.count(), 0);
assert_eq!(wr.count(), 0);
io.st().flags.set_read_paused();
assert!(io.poll_shutdown(&mut cx).is_pending());
assert!(!io.st().flags.is_read_paused());
assert_eq!(rd.count(), 1);
}
#[ntex::test]
async fn filter_shutdown_applies_read_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024 * 1024);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(1024, 256)
.set_shutdown_timeout(ntex_util::time::Seconds(30)),
),
)
.add_filter(StuckShutdown(Cell::new(false)));
let ioref = io.get_ref();
let high = 1024;
ntex::rt::spawn(async move {
let _ = io.shutdown().await;
});
sleep(Millis(50)).await;
for _ in 0..40 {
client.write("A".repeat(1024));
sleep(Millis(5)).await;
}
sleep(Millis(100)).await;
let buffered = ioref.with_read_dst(|buf| buf.len());
assert!(
buffered <= high * 2,
"read buffer grew to {buffered} with a high watermark of {high}"
);
assert!(
client.remote_buffer(|buf| !buf.is_empty()),
"peer send buffer was drained despite read backpressure"
);
}
#[ntex::test]
async fn filter_shutdown_blocked_by_unconsumed_input() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024 * 1024);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(1024, 256)
.set_shutdown_timeout(ntex_util::time::Seconds(30)),
),
)
.add_filter(StuckShutdown(Cell::new(false)));
client.write("A".repeat(4096));
sleep(Millis(50)).await;
assert!(
io.get_ref().is_rd_backpressure(),
"read backpressure was not active before the shutdown"
);
let err = timeout(Millis(3000), io.shutdown())
.await
.expect("shutdown did not complete")
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Other);
}
#[derive(Debug)]
struct Passthrough;
impl FilterLayer for Passthrough {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_read_buffers(|src, dst| {
if let Some(src) = src {
dst.extend_from_slice(src);
src.clear();
}
});
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
}
#[derive(Debug, Default)]
struct AckShutdown {
sent: Cell<bool>,
acked: Cell<bool>,
}
impl FilterLayer for AckShutdown {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_read_buffers(|src, dst| {
if let Some(src) = src {
if self.sent.get() && &src[..] == b"ack" {
self.acked.set(true);
} else {
dst.extend_from_slice(src);
}
src.clear();
}
});
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
if !self.sent.replace(true) {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"bye"));
}
Ok(if self.acked.get() {
Poll::Ready(())
} else {
Poll::Pending
})
}
}
async fn filter_shutdown_waits_for_peer<F: Filter>(client: IoTest, io: Io<F>) {
client.remote_buffer_cap(1024);
let io = io.add_filter(AckShutdown::default());
let done = Rc::new(Cell::new(false));
let done2 = done.clone();
ntex::rt::spawn(async move {
io.shutdown().await.unwrap();
done2.set(true);
});
sleep(Millis(50)).await;
assert_eq!(client.read_any(), Bytes::from_static(b"bye"));
assert!(!done.get());
client.write("ack");
sleep(Millis(50)).await;
assert!(done.get());
}
#[ntex::test]
async fn filter_shutdown_completes_on_peer_input() {
let (client, server) = IoTest::create();
let io = Io::new(server, SharedCfg::new("SRV"));
filter_shutdown_waits_for_peer(client, io).await;
}
#[ntex::test]
async fn filter_shutdown_output_passes_inner_filters() {
let (client, server) = IoTest::create();
let io = Io::new(server, SharedCfg::new("SRV")).add_filter(Passthrough);
filter_shutdown_waits_for_peer(client, io).await;
}
#[ntex::test]
async fn filter_failure_output_passes_inner_filters() {
#[derive(Debug)]
struct FailOnInput;
impl FilterLayer for FailOnInput {
fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
if buf.with_read_src(|src| src.take().is_some()) {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"err"));
buf.io().close();
Err(io::Error::new(io::ErrorKind::InvalidData, "failed"))
} else {
Ok(())
}
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
}
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"))
.add_filter(Passthrough)
.add_filter(FailOnInput);
client.write("input");
let err = io.recv(&BytesCodec).await.unwrap_err();
assert_eq!(err.into_inner().kind(), io::ErrorKind::InvalidData);
sleep(Millis(50)).await;
assert_eq!(client.read_any(), Bytes::from_static(b"err"));
assert!(io.is_closed());
}
#[ntex::test]
async fn filter_shutdown_timeout_is_reported() {
#[derive(Debug)]
struct PendingShutdown;
impl FilterLayer for PendingShutdown {
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<()>> {
Ok(Poll::Pending)
}
}
let (_client, server) = IoTest::create();
let io = Io::new(
server,
SharedCfg::new("SRV")
.add(IoConfig::default().set_shutdown_timeout(ntex_util::time::Seconds(1))),
)
.add_filter(PendingShutdown);
let err = timeout(Millis(3000), io.shutdown())
.await
.expect("transport shutdown did not complete")
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::TimedOut);
assert!(io.st().flags.is_closed());
assert!(!io.st().flags.is_terminating());
}
#[ntex::test]
async fn blocked_filter_shutdown_flushes_buffered_output() {
#[derive(Debug)]
struct ClosingShutdown(Cell<bool>);
impl FilterLayer for ClosingShutdown {
fn process_read_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
if !self.0.replace(true) {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"bye"));
}
Ok(Poll::Pending)
}
}
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(8, 4)
.set_shutdown_timeout(ntex_util::time::Seconds(10)),
),
)
.add_filter(ClosingShutdown(Cell::new(false)));
io.st().flags.set_read_ready_and_backpressure();
io.close();
sleep(Millis(50)).await;
assert!(io.st().flags.is_stopping());
client.remote_buffer_cap(1024);
assert_eq!(
timeout(Millis(1000), client.read())
.await
.expect("closing record was not written")
.unwrap(),
Bytes::from_static(b"bye")
);
let err = timeout(Millis(1000), io.shutdown())
.await
.expect("transport shutdown did not complete")
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Other);
}
#[ntex::test]
async fn one_deadline_bounds_both_shutdown_phases() {
#[derive(Debug)]
struct ClosingShutdown(Cell<bool>);
impl FilterLayer for ClosingShutdown {
fn process_read_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
buf.with_write_buffers(BytePages::move_to);
Ok(())
}
fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
if !self.0.replace(true) {
buf.with_write_buffers(|_, dst| dst.extend_from_slice(b"bye"));
}
Ok(Poll::Pending)
}
}
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(8, 4)
.set_shutdown_timeout(ntex_util::time::Seconds(1)),
),
)
.add_filter(ClosingShutdown(Cell::new(false)));
let start = std::time::Instant::now();
io.close();
timeout(Millis(5000), io.shutdown())
.await
.expect("transport shutdown did not complete")
.unwrap_err();
assert!(io.st().flags.is_closed());
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_millis(1600),
"shutdown took {elapsed:?}, the deadline did not span both phases"
);
}
#[ntex::test]
async fn blocked_filter_shutdown_is_reported() {
#[derive(Debug)]
struct PendingShutdown;
impl FilterLayer for PendingShutdown {
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<()>> {
Ok(Poll::Pending)
}
}
let (_client, server) = IoTest::create();
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_read_buf(8, 4)
.set_shutdown_timeout(ntex_util::time::Seconds(10)),
),
)
.add_filter(PendingShutdown);
io.st().flags.set_read_ready_and_backpressure();
io.close();
sleep(Millis(50)).await;
let err = timeout(Millis(1000), io.shutdown())
.await
.expect("transport shutdown did not complete")
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Other);
assert!(io.st().flags.is_closed());
assert!(!io.st().flags.is_terminating());
}
#[ntex::test]
async fn shutdown() {
#[derive(Debug)]
struct F;
impl FilterLayer for F {
fn process_read_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
fn process_write_buf(&self, _: &FilterBuf<'_>) -> io::Result<()> {
Ok(())
}
}
let io = Io::new(
IoTest::create().0,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(8)),
);
let st = io.st();
assert!(lazy(|cx| io.poll_status_update(cx)).await.is_pending());
assert!(st.dispatch_task.is_set());
assert!(!st.flags.is_peer_gone());
assert!(!st.flags.is_stopping_filters());
let ctx = IoContext::new(io.get_ref());
io.close();
assert!(!st.flags.is_peer_gone());
assert!(st.flags.is_stopping_filters());
let err = io.with_write_src(|_| 1).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Other);
let io = io.add_filter(F);
let layer = Layer::new(F, Base::new(io.get_ref()));
let st = io.st();
st.buffer.with_write_src(|p| p.put_slice(b"123"));
assert_eq!(st.buffer.write_buf_size(), 3);
let res = st.buffer.with_filter(io.as_ref(), |f| layer.shutdown(f));
assert!(matches!(res, Ok(Poll::Ready(()))));
assert_eq!(st.buffer.write_buf_size(), 0);
ctx.stop(None);
assert!(st.flags.is_peer_gone());
assert!(st.flags.is_terminating());
assert!(!st.flags.is_closed());
assert!(st.flags.is_stopping_filters());
let err = io.with_write_src(|_| 1).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::NotConnected);
ctx.stopped(None);
assert!(st.flags.is_closed());
}
struct FixedSize(usize);
impl Decoder for FixedSize {
type Item = Bytes;
type Error = io::Error;
fn decode(&self, src: &mut BytesMut) -> Result<Option<Bytes>, io::Error> {
if src.len() < self.0 {
Ok(None)
} else {
Ok(Some(src.split_to(self.0)))
}
}
}
#[ntex::test]
async fn recv_reports_timeout_during_write_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(16)),
);
io.encode_slice(BIN2).unwrap();
assert!(io.flags().is_wr_backpressure());
let ioref = io.get_ref();
let (res, ()) =
ntex_util::future::join(timeout(Millis(1000), io.recv(&BytesCodec)), async move {
sleep(Millis(25)).await;
ioref.0.notify_timeout();
})
.await;
let Err(Either::Right(err)) = res.expect("recv ignored the timeout") else {
panic!("expected a timeout error")
};
assert_eq!(err.kind(), io::ErrorKind::TimedOut);
assert!(io.flags().is_wr_backpressure());
}
#[ntex::test]
async fn write_timeout_bounds_output_waits() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(
IoConfig::default()
.set_write_buf(16)
.set_write_timeout(ntex_util::time::Seconds(1)),
),
);
io.encode_slice(BIN2).unwrap();
assert!(io.flags().is_wr_backpressure());
let ioref = io.get_ref();
let ((send, flush), ready) = timeout(
Millis(3000),
ntex_util::future::join(
ntex_util::future::join(
io.send(Bytes::from_static(b"item"), &BytesCodec),
io.flush(false),
),
ioref.write_ready(),
),
)
.await
.expect("output waits are not bounded by the write timeout");
let Err(Either::Right(err)) = send else {
panic!("expected a transport error")
};
assert_eq!(err.kind(), io::ErrorKind::TimedOut);
assert_eq!(flush.unwrap_err().kind(), io::ErrorKind::TimedOut);
assert_eq!(ready.unwrap_err().kind(), io::ErrorKind::TimedOut);
assert!(!io.is_closed());
client.remote_buffer_cap(1024);
io.flush(true).await.unwrap();
}
#[ntex::test]
async fn recv_reports_truncated_stream() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("123");
sleep(Millis(25)).await;
client.close().await;
let err = io.recv(&FixedSize(8)).await.err().unwrap();
let Either::Right(err) = err else {
panic!("expected a transport error")
};
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
assert_eq!(io.with_read_dst(|b| b.len()), 3);
}
#[ntex::test]
async fn recv_reports_clean_eof() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("12345678");
sleep(Millis(25)).await;
client.close().await;
assert_eq!(io.recv(&FixedSize(8)).await.unwrap().unwrap(), "12345678");
assert!(io.recv(&FixedSize(8)).await.unwrap().is_none());
}
struct FixedSizeEof(usize, std::cell::Cell<usize>);
impl Decoder for FixedSizeEof {
type Item = Bytes;
type Error = io::Error;
fn decode(&self, src: &mut BytesMut) -> Result<Option<Bytes>, io::Error> {
FixedSize(self.0).decode(src)
}
fn decode_eof(&self, src: &mut BytesMut) -> Result<Option<Bytes>, io::Error> {
self.1.set(self.1.get() + 1);
if src.is_empty() {
Ok(None)
} else {
let len = src.len().min(self.0);
Ok(Some(src.split_to(len)))
}
}
}
#[ntex::test]
async fn recv_uses_decode_eof_after_eof() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
let codec = FixedSizeEof(8, std::cell::Cell::new(0));
client.write("12345678");
sleep(Millis(25)).await;
assert_eq!(io.recv(&codec).await.unwrap().unwrap(), "12345678");
assert_eq!(codec.1.get(), 0);
client.write("123");
sleep(Millis(25)).await;
client.close().await;
assert_eq!(io.recv(&codec).await.unwrap().unwrap(), "123");
assert!(io.recv(&codec).await.unwrap().is_none());
assert!(codec.1.get() >= 2);
assert_eq!(io.with_read_dst(|b| b.len()), 0);
}
#[ntex::test]
async fn recv_local_shutdown_is_not_truncation() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("123");
sleep(Millis(25)).await;
io.close();
sleep(Millis(25)).await;
assert!(io.recv(&FixedSize(8)).await.unwrap().is_none());
assert_eq!(io.with_read_dst(|b| b.len()), 3);
}
#[ntex::test]
async fn read_pause_stops_transport_reads() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
assert!(lazy(|cx| io.poll_read_pause(cx)).await.is_pending());
assert!(io.flags().is_read_paused());
assert!(lazy(|cx| io.poll_read_pause(cx)).await.is_pending());
client.write("data");
sleep(Millis(25)).await;
assert_eq!(io.st().buffer.read_dst_size(), 0);
io.st().notify_timeout();
assert!(matches!(
lazy(|cx| io.poll_read_pause(cx)).await,
Poll::Ready(IoStatusUpdate::Timeout)
));
assert_eq!(io.read_notify().await.unwrap(), Some(()));
assert!(!io.flags().is_read_paused());
assert_eq!(io.with_read_dst(BytesMut::take), b"data");
}
#[ntex::test]
async fn read_dst_size_keeps_read_state() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
assert!(lazy(|cx| io.poll_read_pause(cx)).await.is_pending());
assert_eq!(io.read_dst_size(), 0);
assert!(io.flags().is_read_paused());
client.write("data");
assert_eq!(io.read_notify().await.unwrap(), Some(()));
io.st().flags.set_read_ready();
assert_eq!(io.read_dst_size(), 4);
assert!(io.flags().is_read_ready());
assert_eq!(io.with_read_dst(BytesMut::take), b"data");
assert!(!io.flags().is_read_ready());
}
struct Failing;
impl Decoder for Failing {
type Item = Bytes;
type Error = &'static str;
fn decode(&self, _: &mut BytesMut) -> Result<Option<Bytes>, &'static str> {
Err("invalid frame")
}
}
#[ntex::test]
async fn recv_reports_decoder_error() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("data");
let Err(Either::Left(err)) = io.recv(&Failing).await else {
panic!("expected a decoder error")
};
assert_eq!(err, "invalid frame");
}
#[ntex::test]
async fn poll_recv_decode_reports_timeout_before_decoding() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("data");
sleep(Millis(25)).await;
io.st().notify_timeout();
let res = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx)).await;
assert!(matches!(res, Err(RecvError::Timeout)));
assert_eq!(io.st().buffer.read_dst_size(), 4);
let decoded = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx))
.await
.unwrap();
assert_eq!(decoded.item.unwrap(), "data");
assert_eq!((decoded.consumed, decoded.remains), (4, 0));
}
#[ntex::test]
async fn poll_recv_decode_reports_write_backpressure_before_decoding() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("data");
sleep(Millis(25)).await;
io.st().flags.set_wr_backpressure();
let res = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx)).await;
assert!(matches!(res, Err(RecvError::WriteBackpressure)));
assert_eq!(io.st().buffer.read_dst_size(), 4);
io.flush(false).await.unwrap();
let decoded = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx))
.await
.unwrap();
assert_eq!(decoded.item.unwrap(), "data");
}
#[ntex::test]
async fn poll_recv_decode_closing_decodes_buffered_input() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(1024);
let io = Io::new(server, SharedCfg::new("SRV"));
client.write("data");
sleep(Millis(25)).await;
io.close();
sleep(Millis(25)).await;
io.st().notify_timeout();
io.st().flags.set_wr_backpressure();
let decoded = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx))
.await
.unwrap();
assert_eq!(decoded.item.unwrap(), "data");
let res = lazy(|cx| io.poll_recv_decode(&BytesCodec, cx)).await;
assert!(matches!(res, Err(RecvError::PeerGone(None))));
}
#[ntex::test]
async fn poll_flush_enables_write_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let io = Io::new(
server,
SharedCfg::new("SRV").add(IoConfig::default().set_write_buf(16)),
);
io.with_write_dst(|buf| buf.extend_from_slice(BIN2));
assert!(!io.flags().is_wr_backpressure());
assert!(lazy(|cx| io.poll_flush(cx, false)).await.is_pending());
assert!(io.flags().is_wr_backpressure());
client.remote_buffer_cap(1024);
assert_eq!(client.read().await.unwrap(), BIN2);
io.flush(false).await.unwrap();
assert!(!io.flags().is_wr_backpressure());
}
struct Gate<F>(F, Rc<Cell<bool>>, Rc<Cell<bool>>);
impl<F: Filter> Filter for Gate<F> {
fn query(&self, id: std::any::TypeId) -> Option<Box<dyn std::any::Any>> {
self.0.query(id)
}
fn process_read_buf(&self, ctx: &mut crate::FilterCtx<'_>) -> io::Result<()> {
self.0.process_read_buf(ctx)
}
fn process_write_buf(&self, ctx: &mut crate::FilterCtx<'_>) -> io::Result<()> {
self.0.process_write_buf(ctx)
}
fn shutdown(&self, ctx: &mut crate::FilterCtx<'_>) -> io::Result<Poll<()>> {
self.0.shutdown(ctx)
}
fn poll_read_ready(&self, cx: &mut Context<'_>) -> Poll<Readiness> {
match self.0.poll_read_ready(cx) {
Poll::Ready(Readiness::Ready) if self.1.get() => Poll::Pending,
res => res,
}
}
fn poll_write_ready(&self, cx: &mut Context<'_>) -> Poll<Readiness> {
match self.0.poll_write_ready(cx) {
Poll::Ready(Readiness::Ready) if self.2.get() => Poll::Pending,
res => res,
}
}
}
#[derive(Debug)]
struct GateTransport;
impl IoStream for GateTransport {
fn start(self, _: IoContext) -> Box<dyn Handle> {
Box::new(self)
}
}
impl Handle for GateTransport {}
#[ntex::test]
async fn filter_pause_pauses_reading() {
let blocked = Rc::new(Cell::new(false));
let b = blocked.clone();
let io = Io::new(GateTransport, SharedCfg::default())
.map_filter(move |f| Gate(f, b, Rc::default()));
let ctx = IoContext::new(io.get_ref());
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.is_read_filter_paused());
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
blocked.set(true);
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(io.is_read_filter_paused());
assert!(!io.st().dispatch_task.is_set(), "dispatcher is notified");
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(io.st().dispatch_task.is_set());
let mut buf = ctx.take_read_buf();
buf.extend_from_slice(b"12");
assert_eq!(
ctx.release_read_buf(buf, Poll::Ready(Ok(2))),
IoTaskStatus::Pause
);
assert_eq!(io.with_read_dst(BytesMut::take), b"12");
assert_eq!(
ctx.with_read_buf(|_| Poll::<io::Result<usize>>::Pending),
IoTaskStatus::Pause
);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
blocked.set(false);
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.is_read_filter_paused());
assert!(!io.st().dispatch_task.is_set(), "dispatcher is notified");
assert_eq!(
ctx.with_read_buf(|_| Poll::<io::Result<usize>>::Pending),
IoTaskStatus::Io
);
}
#[ntex::test]
async fn io_state_pause_is_not_filter_pause() {
let blocked = Rc::new(Cell::new(true));
let b = blocked.clone();
let io = Io::new(GateTransport, SharedCfg::default())
.map_filter(move |f| Gate(f, b, Rc::default()));
let ctx = IoContext::new(io.get_ref());
io.st().flags.set_read_paused();
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(!io.is_read_filter_paused());
io.st().flags.unset_read_paused();
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(io.is_read_filter_paused());
io.st().flags.set_read_paused();
assert!(lazy(|cx| ctx.poll_read_ready(cx)).await.is_pending());
assert!(io.is_read_filter_paused());
blocked.set(false);
io.st().flags.unset_read_paused();
assert_eq!(
lazy(|cx| ctx.poll_read_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.is_read_filter_paused());
}
#[ntex::test]
async fn filter_pause_pauses_writing() {
let blocked = Rc::new(Cell::new(true));
let b = blocked.clone();
let io = Io::new(GateTransport, SharedCfg::default())
.map_filter(move |f| Gate(f, Rc::default(), b));
let ctx = IoContext::new(io.get_ref());
assert!(lazy(|cx| ctx.poll_write_ready(cx)).await.is_pending());
assert!(!io.is_write_filter_paused());
io.encode_slice(b"1234").unwrap();
sleep(Millis(10)).await;
assert!(!io.st().flags.is_write_paused());
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
assert!(lazy(|cx| ctx.poll_write_ready(cx)).await.is_pending());
assert!(io.is_write_filter_paused());
assert!(!io.st().dispatch_task.is_set(), "dispatcher is notified");
ctx.with_write_dst(|dst| {
let mut page = dst.take().unwrap();
page.advance_to(2);
dst.prepend(page);
});
assert_eq!(ctx.update_write_status(Ok(2)), IoTaskStatus::Pause);
assert!(!io.st().flags.is_write_paused());
assert_eq!(io.st().buffer.write_buf_size(), 2);
assert!(lazy(|cx| io.poll_read_more(cx)).await.is_pending());
blocked.set(false);
assert_eq!(
lazy(|cx| ctx.poll_write_ready(cx)).await,
Poll::Ready(Readiness::Ready)
);
assert!(!io.is_write_filter_paused());
assert!(!io.st().dispatch_task.is_set(), "dispatcher is notified");
assert_eq!(ctx.update_write_status(Ok(0)), IoTaskStatus::Io);
}
}