use std::{cell::Cell, fmt, io, task::Poll};
use ntex_bytes::{BytePageSize, BytePages, BytesMut};
use crate::{IoConfig, IoRef};
pub(crate) struct Stack {
buffers: Vec<Buffer>,
}
#[derive(Default)]
struct Buffer {
read: Cell<Option<BytesMut>>,
write: Cell<Option<BytePages>>,
}
impl fmt::Debug for Stack {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Stack")
.field("len", &self.buffers.len())
.finish()
}
}
impl Stack {
pub(crate) fn new(size: BytePageSize) -> Self {
let mut buffers = Vec::with_capacity(4);
buffers.push(Buffer {
read: Cell::new(None),
write: Cell::new(Some(BytePages::new(size))),
});
buffers.push(Buffer::default());
Self { buffers }
}
pub(crate) fn set_page_size(&self, size: BytePageSize) {
for b in &self.buffers {
b.with_write_if_set(|b| b.set_page_size(size));
}
}
pub(crate) fn add_layer(&mut self, page_size: BytePageSize) {
self.buffers.insert(
0,
Buffer {
read: Cell::new(None),
write: Cell::new(Some(BytePages::new(page_size))),
},
);
}
fn with_last<F, R>(&self, f: F) -> R
where
F: FnOnce(&Buffer) -> R,
{
f(&self.buffers[self.buffers.len() - 2])
}
pub(crate) fn with_read_src<F, R>(&self, io: &IoRef, f: F) -> R
where
F: FnOnce(&mut BytesMut) -> R,
{
self.with_last(|buf| buf.with_read(io, f))
}
pub(crate) fn with_read_dst<F, R>(&self, io: &IoRef, f: F) -> R
where
F: FnOnce(&mut BytesMut) -> R,
{
self.buffers[0].with_read(io, f)
}
pub(crate) fn write_buf_size(&self) -> usize {
if self.buffers.len() == 2 {
self.buffers[0].write_len()
} else {
self.buffers[0].write_len() + self.buffers[self.buffers.len() - 2].write_len()
}
}
pub(crate) fn write_dst_size(&self) -> usize {
self.with_last(Buffer::write_len)
}
pub(crate) fn with_write_src<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut BytePages) -> R,
{
self.buffers[0].with_write(f)
}
pub(crate) fn with_write_dst<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut BytePages) -> R,
{
self.buffers[self.buffers.len() - 2].with_write(f)
}
pub(crate) fn read_dst_size(&self) -> usize {
self.buffers[0].read_len()
}
pub(crate) fn with_filter<F, R>(&self, io: &IoRef, f: F) -> R
where
F: FnOnce(&mut FilterCtx<'_>) -> R,
{
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: false },
};
f(&mut ctx)
}
pub(crate) fn get_read_buf(&self) -> Option<BytesMut> {
self.with_last(|buffer| buffer.read.take())
}
pub(crate) fn set_read_buf(&self, buf: BytesMut, cfg: &IoConfig) {
self.with_last(move |buffer| {
if let Some(mut first_buf) = buffer.read.take() {
cfg.read_buf().resize_min(&mut first_buf, buf.len());
first_buf.extend_from_slice(&buf);
cfg.read_buf().release(buf);
buffer.read.set(Some(first_buf));
} else if !buf.is_empty() {
buffer.read.set(Some(buf));
} else {
cfg.read_buf().release(buf);
}
});
}
pub(crate) fn process_read_buf(&self, io: &IoRef) -> io::Result<FilterUpdates> {
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: false },
};
io.with_callbacks(|cb| cb.before_processing(io));
let result = io.filter().process_read_buf(&mut ctx);
io.with_callbacks(|cb| cb.after_processing(io));
result.map(|()| ctx.st)
}
pub(crate) fn process_read_buf_no_cb(&self, io: &IoRef) -> io::Result<FilterUpdates> {
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: false },
};
io.filter().process_read_buf(&mut ctx).map(|()| ctx.st)
}
pub(crate) fn process_write_buf(&self, io: &IoRef) -> io::Result<()> {
if self.buffers[0].is_write_empty() {
Ok(())
} else {
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: true },
};
io.with_callbacks(|cb| cb.before_processing(io));
let res = io.filter().process_write_buf(&mut ctx);
io.with_callbacks(|cb| cb.after_processing(io));
res
}
}
pub(crate) fn process_write_buf_no_cb(&self, io: &IoRef) -> io::Result<()> {
if self.buffers[0].is_write_empty() {
Ok(())
} else {
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: true },
};
io.filter().process_write_buf(&mut ctx)
}
}
pub(crate) fn process_write_buf_force(&self, io: &IoRef) -> io::Result<()> {
let mut ctx = FilterCtx {
io,
idx: 0,
stack: self,
st: FilterUpdates { wants_write: true },
};
io.with_callbacks(|cb| cb.before_processing(io));
let res = io.filter().process_write_buf(&mut ctx);
io.with_callbacks(|cb| cb.after_processing(io));
res
}
pub(crate) fn process_shutdown(&self, io: &IoRef) -> io::Result<Poll<()>> {
self.process_write_buf(io)?;
io.with_callbacks(|cb| cb.before_processing(io));
let res = self.with_filter(io, |ctx| io.filter().shutdown(ctx));
io.with_callbacks(|cb| cb.after_processing(io));
res
}
pub(crate) fn release(&self, cfg: &IoConfig) {
for b in &self.buffers {
if let Some(buf) = b.read.take() {
cfg.read_buf().release(buf);
}
b.with_write_if_set(BytePages::clear);
}
}
}
impl Buffer {
fn with_write_if_set(&self, f: impl FnOnce(&mut BytePages)) {
if let Some(mut wb) = self.write.take() {
f(&mut wb);
self.write.set(Some(wb));
}
}
fn is_write_empty(&self) -> bool {
self.with_write(|b| b.is_empty())
}
fn read_len(&self) -> usize {
if let Some(rb) = self.read.take() {
let l = rb.len();
self.read.set(Some(rb));
l
} else {
0
}
}
fn write_len(&self) -> usize {
self.with_write(|b| b.len())
}
fn with_read<F, R>(&self, io: &IoRef, f: F) -> R
where
F: FnOnce(&mut BytesMut) -> R,
{
let mut rb = self
.read
.take()
.unwrap_or_else(|| io.cfg().read_buf().get());
let result = f(&mut rb);
if self.read.take().is_some() {
log::error!("Nested read io operation is detected");
io.terminate();
}
if rb.is_empty() {
io.cfg().read_buf().release(rb);
} else {
self.read.set(Some(rb));
}
result
}
fn with_write<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut BytePages) -> R,
{
let mut wb = self.write.take().unwrap();
let result = f(&mut wb);
self.write.set(Some(wb));
result
}
}
#[derive(Copy, Clone, Debug)]
pub(crate) struct FilterUpdates {
pub(crate) wants_write: bool,
}
#[derive(Debug)]
pub struct FilterCtx<'a> {
io: &'a IoRef,
idx: usize,
stack: &'a Stack,
st: FilterUpdates,
}
impl FilterCtx<'_> {
#[inline]
pub fn io(&self) -> &IoRef {
self.io
}
#[inline]
pub fn tag(&self) -> &'static str {
self.io.tag()
}
#[inline]
pub fn with_next<F, R>(&mut self, f: F) -> R
where
F: FnOnce(&mut Self) -> R,
{
self.idx += 1;
let res = f(self);
self.idx -= 1;
res
}
#[inline]
pub fn with_buffer<F, R>(&mut self, f: F) -> R
where
F: FnOnce(&mut FilterBuf<'_>) -> R,
{
let mut buf = FilterBuf {
io: self.io,
curr: &self.stack.buffers[self.idx],
next: &self.stack.buffers[self.idx + 1],
wants_write: Cell::new(self.st.wants_write),
};
let result = f(&mut buf);
if buf.wants_write.get() {
self.st.wants_write = true;
}
result
}
#[inline]
pub fn read_dst_size(&self) -> usize {
self.stack.buffers[0].read_len()
}
#[inline]
pub fn write_dst_size(&mut self) -> usize {
self.stack.buffers[self.stack.buffers.len() - 2].write_len()
}
pub(crate) fn clear_write_buf(&mut self) {
self.stack.buffers[self.idx].with_write(BytePages::clear);
}
}
#[derive(Debug)]
pub struct FilterBuf<'a> {
io: &'a IoRef,
curr: &'a Buffer,
next: &'a Buffer,
wants_write: Cell<bool>,
}
impl FilterBuf<'_> {
#[inline]
pub fn io(&self) -> &IoRef {
self.io
}
#[inline]
pub fn tag(&self) -> &'static str {
self.io.tag()
}
pub fn with_read_src<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut Option<BytesMut>) -> R,
{
let mut read_src = self.next.read.take();
let result = f(&mut read_src);
if let Some(b) = read_src {
if b.is_empty() {
self.io.cfg().read_buf().release(b);
} else {
self.next.read.set(Some(b));
}
}
result
}
pub fn with_read_buffers<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut Option<BytesMut>, &mut BytesMut) -> R,
{
let mut read_src = self.next.read.take();
let mut read_dst = self
.curr
.read
.take()
.unwrap_or_else(|| self.io.cfg().read_buf().get());
let result = f(&mut read_src, &mut read_dst);
if let Some(b) = read_src {
if b.is_empty() {
self.io.cfg().read_buf().release(b);
} else {
self.next.read.set(Some(b));
}
}
if read_dst.is_empty() {
self.io.cfg().read_buf().release(read_dst);
} else {
self.curr.read.set(Some(read_dst));
}
result
}
#[inline]
pub fn with_write_buffers<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut BytePages, &mut BytePages) -> R,
{
let mut write_curr = self.curr.write.take().unwrap();
let (mut write_next, on_demand) = match self.next.write.take() {
Some(b) => (b, false),
None => (BytePages::new(write_curr.page_size()), true),
};
let write_len = if self.wants_write.get() {
0
} else {
write_next.len()
};
let result = f(&mut write_curr, &mut write_next);
if !self.wants_write.get() && write_next.len() > write_len {
self.wants_write.set(true);
}
let undelivered = on_demand && !write_next.is_empty();
self.curr.write.set(Some(write_curr));
if !on_demand || undelivered {
self.next.write.set(Some(write_next));
}
debug_assert!(
!undelivered,
"{}: output written to the write destination of the innermost filter buffer is never sent",
self.io.tag()
);
result
}
}
impl fmt::Debug for Buffer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let read = self.read.take();
let write = self.write.take();
let result = f
.debug_struct("Buffer")
.field("read", &read)
.field("write", &write)
.finish();
self.read.set(read);
self.write.set(write);
result
}
}
#[cfg(test)]
mod tests {
use ntex_bytes::BufMut;
use super::*;
use crate::{Io, testing::IoTest};
#[ntex::test]
async fn stack_buffers() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let mut stack = Stack::new(BytePageSize::Size8);
assert!(format!("{stack:?}").contains("len: 2"));
assert_eq!(stack.read_dst_size(), 0);
assert_eq!(stack.write_buf_size(), 0);
let ptr = stack.buffers.as_ptr();
stack.add_layer(BytePageSize::Size16);
stack.add_layer(BytePageSize::Size16);
assert_eq!(stack.buffers.as_ptr(), ptr);
stack.buffers.remove(0);
assert_eq!(stack.buffers.len(), 3);
stack.set_page_size(BytePageSize::Size32);
let (inner, layers) = stack.buffers.split_last().unwrap();
for buffer in layers {
buffer.with_write(|buf| assert_eq!(buf.page_size(), BytePageSize::Size32));
}
assert!(inner.write.take().is_none());
stack.set_read_buf(BytesMut::from(&b"one"[..]), ioref.cfg());
stack.set_read_buf(BytesMut::from(&b"-two"[..]), ioref.cfg());
assert_eq!(stack.get_read_buf().as_deref(), Some(b"one-two".as_ref()));
assert!(stack.get_read_buf().is_none());
stack.set_read_buf(BytesMut::new(), ioref.cfg());
assert!(stack.get_read_buf().is_none());
}
#[ntex::test]
async fn set_read_buf_merges_into_cacheable_buffer() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let cfg = ioref.cfg().read_buf();
let stack = Stack::new(BytePageSize::Size8);
let mut first = cfg.get();
first.extend_from_slice(&vec![1; cfg.high - 100]);
let frame = first.split_to(cfg.high - 1100);
stack.set_read_buf(first, ioref.cfg());
let mut second = cfg.get();
second.extend_from_slice(&[2; 4000]);
stack.set_read_buf(second, ioref.cfg());
let merged = stack.get_read_buf().unwrap();
assert_eq!(merged.len(), 5000);
assert_eq!(&merged[..1000], &[1; 1000][..]);
assert_eq!(&merged[1000..], &[2; 4000][..]);
assert_eq!(merged.capacity(), cfg.high);
assert_eq!(frame.len(), cfg.high - 1100);
}
#[ntex::test]
async fn filter_read_buffers() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let mut stack = Stack::new(BytePageSize::Size8);
stack.add_layer(BytePageSize::Size8);
stack.set_read_buf(BytesMut::from(&b"input"[..]), ioref.cfg());
stack.with_filter(&ioref, |ctx| {
assert_eq!(ctx.io(), &ioref);
assert_eq!(ctx.tag(), ioref.tag());
assert_eq!(ctx.read_dst_size(), 0);
ctx.with_buffer(|buf| {
assert_eq!(buf.io(), &ioref);
assert_eq!(buf.tag(), ioref.tag());
buf.with_read_buffers(|src, dst| {
let src = src.as_mut().unwrap();
dst.extend_from_slice(&src.split_to(2));
});
});
});
assert_eq!(stack.read_dst_size(), 2);
stack.with_read_dst(&ioref, |buf| assert_eq!(&buf[..], b"in"));
assert_eq!(stack.get_read_buf().as_deref(), Some(b"put".as_ref()));
stack.with_filter(&ioref, |ctx| {
ctx.with_buffer(|buf| {
buf.with_read_src(|src| {
*src = Some(BytesMut::from(&b"next"[..]));
});
});
});
assert_eq!(stack.get_read_buf().as_deref(), Some(b"next".as_ref()));
}
#[ntex::test]
async fn innermost_write_buffer_is_on_demand() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let stack = Stack::new(BytePageSize::Size8);
let inner = &stack.buffers[1];
assert!(inner.write.take().is_none());
stack.release(ioref.cfg());
stack.with_write_src(|buf| buf.put_slice(b"out"));
stack.with_filter(&ioref, |ctx| {
ctx.with_buffer(|buf| {
buf.with_write_buffers(|src, dst| {
assert_eq!(src.len(), 3);
assert!(dst.is_empty());
});
});
});
assert!(inner.write.take().is_none());
}
#[cfg(debug_assertions)]
#[ntex::test]
#[should_panic(expected = "is never sent")]
async fn innermost_write_destination_output_asserts() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let stack = Stack::new(BytePageSize::Size8);
stack.with_write_src(|buf| buf.put_slice(b"out"));
stack.with_filter(&ioref, |ctx| {
ctx.with_buffer(|buf| buf.with_write_buffers(BytePages::move_to));
});
}
#[ntex::test]
async fn filter_write_buffers_and_updates() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let mut stack = Stack::new(BytePageSize::Size8);
stack.add_layer(BytePageSize::Size8);
stack.with_write_src(|buf| buf.put_slice(b"output"));
let updates = stack.with_filter(&ioref, |ctx| {
assert_eq!(ctx.write_dst_size(), 0);
ctx.with_buffer(|buf| {
buf.with_write_buffers(|src, dst| {
assert_eq!(src.len(), 6);
assert_eq!(dst.len(), 0);
src.move_to(dst);
});
});
ctx.st
});
assert!(updates.wants_write);
assert_eq!(stack.write_buf_size(), 6);
assert_eq!(
stack.with_write_dst(|buf| buf.split_to(6).freeze()),
b"output".as_ref()
);
assert_eq!(stack.write_buf_size(), 0);
}
#[ntex::test]
async fn buffer_debug_preserves_contents() {
let (_, server) = IoTest::create();
let io = Io::from(server);
let ioref = io.get_ref();
let buffer = Buffer {
read: Cell::new(None),
write: Cell::new(Some(BytePages::new(BytePageSize::Size8))),
};
buffer.with_read(&ioref, |buf| buf.extend_from_slice(b"read"));
buffer.with_write(|buf| buf.put_slice(b"write"));
let debug = format!("{buffer:?}");
assert!(debug.contains("Buffer"));
assert_eq!(buffer.read_len(), 4);
assert_eq!(buffer.write_len(), 5);
}
}