use std::cell::Cell;
use std::io;
use std::task::{Context, Poll, Waker};
use ntex_service::state::{RequestState, State};
use ntex_util::time::{Seconds, Sleep};
use crate::waiters::{WaiterEntry, Waiters};
use crate::{Filter, Io, IoBoxed, IoCallbacks};
#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct Decoded<T> {
pub item: Option<T>,
pub remains: usize,
pub consumed: usize,
}
pub(crate) struct Extensions(Cell<Option<Box<ExtensionsInner>>>);
#[derive(Default)]
pub(crate) struct ExtensionsInner {
waiters: Waiters,
pub(crate) callbacks: Option<Box<dyn IoCallbacks>>,
}
impl Default for Extensions {
fn default() -> Extensions {
Extensions(Cell::new(None))
}
}
impl Extensions {
fn with<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut ExtensionsInner) -> R,
{
let mut inner = if let Some(inner) = self.0.take() {
inner
} else {
Box::new(ExtensionsInner::default())
};
let result = f(&mut inner);
self.0.set(Some(inner));
result
}
fn with_opt<F>(&self, f: F)
where
F: FnOnce(&mut ExtensionsInner),
{
if let Some(mut inner) = self.0.take() {
f(&mut inner);
self.0.set(Some(inner));
}
}
pub(super) fn register_waker(&self, waiter: &WaiterEntry, waker: &Waker) {
self.with(|inner| {
let waiters = &mut inner.waiters;
if !waiter.id.get().is_some_and(|id| waiters.update(id, waker)) {
waiter.id.set(Some(waiters.register(waiter.tag, waker)));
}
});
}
pub(super) fn poll_waker(&self, waiter: &WaiterEntry, waker: &Waker) -> Poll<()> {
self.with(|inner| {
let waiters = &mut inner.waiters;
match waiter.id.get() {
None => {
waiter.id.set(Some(waiters.register(waiter.tag, waker)));
Poll::Pending
}
Some(id) if waiters.update(id, waker) => Poll::Pending,
Some(_) => {
waiter.id.set(None);
Poll::Ready(())
}
}
})
}
pub(super) fn remove_waker(&self, waiter: &WaiterEntry) {
if let Some(id) = waiter.id.take() {
self.with_opt(|inner| inner.waiters.remove(id, waiter.tag));
}
}
pub(super) fn wake(&self, tag: usize) {
self.with_opt(|inner| inner.waiters.wake(tag));
}
pub(super) fn wake_all(&self) {
self.with_opt(|inner| inner.waiters.wake_all());
}
#[cfg(test)]
pub(super) fn wakers_len(&self) -> usize {
let mut len = 0;
self.with_opt(|inner| len = inner.waiters.len());
len
}
pub(super) fn register_filter_callbacks<T: IoCallbacks + 'static>(&self, cb: T) {
self.with(|inner| {
inner.callbacks = Some(Box::new(cb));
});
}
pub(super) fn take_callbacks(&self) -> Option<Box<dyn IoCallbacks>> {
let mut callbacks = None;
self.with_opt(|inner| callbacks = inner.callbacks.take());
callbacks
}
pub(crate) fn with_callbacks<F>(&self, f: F)
where
F: FnOnce(&dyn IoCallbacks),
{
self.with_opt(|inner| {
if let Some(ref cb) = inner.callbacks {
f(cb.as_ref());
}
});
}
}
impl<F> RequestState<Io<F>> for Io<F> {
type State = ();
#[inline]
fn unpack(self) -> ((), Io<F>) {
((), self)
}
}
impl<F: Filter> RequestState<IoBoxed> for Io<F> {
type State = ();
#[inline]
fn unpack(self) -> ((), IoBoxed) {
((), self.boxed())
}
}
impl RequestState<IoBoxed> for IoBoxed {
type State = ();
#[inline]
fn unpack(self) -> ((), IoBoxed) {
((), self)
}
}
impl<F: Filter, St: 'static> RequestState<IoBoxed> for State<St, Io<F>> {
type State = St;
#[inline]
fn unpack(self) -> (St, IoBoxed) {
let State { req, state } = self;
(state, req.boxed())
}
}
pub(crate) struct WriteDeadline {
timeout: Seconds,
sleep: Option<Sleep>,
}
impl WriteDeadline {
pub(crate) fn new(timeout: Seconds) -> Self {
Self {
timeout,
sleep: None,
}
}
pub(crate) fn poll_expired(&mut self, cx: &mut Context<'_>) -> bool {
if self.timeout.is_zero() {
false
} else {
let timeout = self.timeout;
self.sleep
.get_or_insert_with(|| Sleep::new(timeout.into()))
.poll_elapsed(cx)
.is_ready()
}
}
}
pub(crate) fn write_timed_out() -> io::Error {
io::Error::new(io::ErrorKind::TimedOut, "Write timeout")
}
#[cfg(test)]
mod tests {
use ntex_bytes::BytePageSize;
use ntex_service::cfg::SharedCfg;
use super::*;
use crate::{Sealed, buf::Stack, filter::NullFilter, testing::IoTest};
#[ntex::test]
async fn test_null_filter() {
let (_, server) = IoTest::create();
let io = Io::new(server, SharedCfg::default());
let ioref = io.get_ref();
let stack = Stack::new(BytePageSize::Size16);
assert!(NullFilter.query(std::any::TypeId::of::<()>()).is_none());
assert!(
stack
.with_filter(&ioref, |ctx| NullFilter.shutdown(ctx))
.unwrap()
.is_ready()
);
assert_eq!(
std::future::poll_fn(|cx| NullFilter.poll_read_ready(cx)).await,
crate::Readiness::Close
);
assert_eq!(
std::future::poll_fn(|cx| NullFilter.poll_write_ready(cx)).await,
crate::Readiness::Close
);
assert!(
stack
.with_filter(&ioref, |ctx| NullFilter.process_write_buf(ctx))
.is_ok()
);
assert_eq!(
stack.with_filter(&ioref, |ctx| NullFilter.process_read_buf(ctx).unwrap()),
()
);
}
#[ntex::test]
async fn request_state_unpack() {
use ntex_service::state::{RequestState, State};
let (_, server) = IoTest::create();
let io = Io::from(server);
let id = io.id();
let ((), io) = <Io as RequestState<Io>>::unpack(io);
assert_eq!(io.id(), id);
let ((), io) = <Io as RequestState<IoBoxed>>::unpack(io);
assert_eq!(io.id(), id);
let ((), mut io) = <IoBoxed as RequestState<IoBoxed>>::unpack(io);
assert_eq!(io.id(), id);
let st = State {
req: Io::<Sealed>::from(io.take()),
state: 10u32,
};
let (state, io) = <_ as RequestState<IoBoxed>>::unpack(st);
assert_eq!(state, 10);
assert_eq!(io.id(), id);
}
}