Skip to main content

ntex_io/
testing.rs

1//! In-memory I/O transport and helpers for tests.
2#![allow(clippy::missing_panics_doc)]
3use std::sync::{Arc, Mutex};
4use std::task::{Context, Poll, Waker};
5use std::{any, cell::RefCell, cmp, fmt, future::poll_fn, io, mem, net, rc::Rc};
6
7use ntex_bytes::{BufMut, BytePages, Bytes, BytesMut};
8use ntex_util::time::{Millis, sleep};
9
10use crate::{Handle, IoContext, IoStream, IoTaskStatus, Readiness, types};
11
12#[derive(Default)]
13struct AtomicWaker(Arc<Mutex<RefCell<Option<Waker>>>>);
14
15impl AtomicWaker {
16    fn wake(&self) -> bool {
17        if let Some(waker) = self.0.lock().unwrap().borrow_mut().take() {
18            waker.wake();
19            true
20        } else {
21            false
22        }
23    }
24}
25
26impl fmt::Debug for AtomicWaker {
27    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
28        write!(f, "AtomicWaker")
29    }
30}
31
32/// One endpoint of an in-memory asynchronous byte stream.
33///
34/// Use [`IoTest::create`] to construct connected client and server endpoints.
35/// Bytes written by one endpoint become readable from the other.
36#[derive(Debug)]
37pub struct IoTest {
38    tp: Type,
39    peer_addr: Option<net::SocketAddr>,
40    state: Arc<Mutex<RefCell<State>>>,
41    local: Arc<Mutex<RefCell<Channel>>>,
42    remote: Arc<Mutex<RefCell<Channel>>>,
43}
44
45bitflags::bitflags! {
46    #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
47    struct IoTestFlags: u8 {
48        const FLUSHED = 0b0000_0001;
49        const CLOSED  = 0b0000_0010;
50    }
51}
52
53#[derive(Copy, Clone, PartialEq, Eq, Debug)]
54enum Type {
55    Client,
56    Server,
57    ClientClone,
58    ServerClone,
59}
60
61#[derive(Copy, Clone, Default, Debug)]
62struct State {
63    client_dropped: bool,
64    server_dropped: bool,
65}
66
67#[derive(Default, Debug)]
68struct Channel {
69    buf: BytesMut,
70    buf_cap: usize,
71    flags: IoTestFlags,
72    waker: AtomicWaker,
73    read: IoTestState,
74    write: IoTestState,
75}
76
77unsafe impl Sync for Channel {}
78unsafe impl Send for Channel {}
79
80impl Channel {
81    fn is_closed(&self) -> bool {
82        self.flags.contains(IoTestFlags::CLOSED)
83    }
84}
85
86impl Default for IoTestFlags {
87    fn default() -> Self {
88        IoTestFlags::empty()
89    }
90}
91
92#[derive(Debug, Default)]
93enum IoTestState {
94    #[default]
95    Ok,
96    Pending,
97    Close,
98    Err(io::Error),
99}
100
101impl IoTest {
102    /// Creates connected client and server endpoints.
103    pub fn create() -> (IoTest, IoTest) {
104        let local = Arc::new(Mutex::new(RefCell::new(Channel::default())));
105        let remote = Arc::new(Mutex::new(RefCell::new(Channel::default())));
106        let state = Arc::new(Mutex::new(RefCell::new(State::default())));
107
108        (
109            IoTest {
110                tp: Type::Client,
111                peer_addr: None,
112                local: local.clone(),
113                remote: remote.clone(),
114                state: state.clone(),
115            },
116            IoTest {
117                state,
118                peer_addr: None,
119                tp: Type::Server,
120                local: remote,
121                remote: local,
122            },
123        )
124    }
125
126    /// Returns `true` after the client endpoint has been dropped.
127    pub fn is_client_dropped(&self) -> bool {
128        self.state.lock().unwrap().borrow().client_dropped
129    }
130
131    /// Returns `true` after the server endpoint has been dropped.
132    pub fn is_server_dropped(&self) -> bool {
133        self.state.lock().unwrap().borrow().server_dropped
134    }
135
136    /// Returns `true` after the peer has closed its write side.
137    pub fn is_closed(&self) -> bool {
138        self.remote.lock().unwrap().borrow().is_closed()
139    }
140
141    /// Sets the socket address returned by transport queries.
142    #[must_use]
143    pub fn set_peer_addr(mut self, addr: net::SocketAddr) -> Self {
144        self.peer_addr = Some(addr);
145        self
146    }
147
148    /// Forces subsequent reads from the peer endpoint to remain pending.
149    pub fn read_pending(&self) {
150        self.remote.lock().unwrap().borrow_mut().read = IoTestState::Pending;
151    }
152
153    /// Makes the peer endpoint's next read fail with `err`.
154    pub fn read_error(&self, err: io::Error) {
155        let channel = self.remote.lock().unwrap();
156        channel.borrow_mut().read = IoTestState::Err(err);
157        channel.borrow().waker.wake();
158    }
159
160    /// Makes the peer endpoint's next transport write fail with `err`.
161    pub fn write_error(&self, err: io::Error) {
162        self.local.lock().unwrap().borrow_mut().write = IoTestState::Err(err);
163        self.remote.lock().unwrap().borrow().waker.wake();
164    }
165
166    /// Makes the peer endpoint's next transport write return zero bytes.
167    pub fn write_zero(&self) {
168        self.local.lock().unwrap().borrow_mut().write = IoTestState::Close;
169        self.remote.lock().unwrap().borrow().waker.wake();
170    }
171
172    /// Provides mutable access to bytes readable by this endpoint.
173    pub fn local_buffer<F, R>(&self, f: F) -> R
174    where
175        F: FnOnce(&mut BytesMut) -> R,
176    {
177        let guard = self.local.lock().unwrap();
178        let mut ch = guard.borrow_mut();
179        f(&mut ch.buf)
180    }
181
182    /// Provides mutable access to bytes readable by the peer endpoint.
183    pub fn remote_buffer<F, R>(&self, f: F) -> R
184    where
185        F: FnOnce(&mut BytesMut) -> R,
186    {
187        let guard = self.remote.lock().unwrap();
188        let mut ch = guard.borrow_mut();
189        f(&mut ch.buf)
190    }
191
192    /// Simulates this endpoint closing its write side.
193    ///
194    /// This wakes the peer endpoint's pending reader and waits briefly so the
195    /// close can be observed by asynchronous test code.
196    pub async fn close(&self) {
197        {
198            let guard = self.remote.lock().unwrap();
199            let mut remote = guard.borrow_mut();
200            remote.read = IoTestState::Close;
201            let res = remote.waker.wake();
202            log::debug!("close remote socket, waker:{res}");
203        }
204        sleep(Millis(35)).await;
205    }
206
207    /// Writes bytes for the peer endpoint to read and wakes its reader.
208    pub fn write<T: AsRef<[u8]>>(&self, data: T) {
209        let guard = self.remote.lock().unwrap();
210        let mut write = guard.borrow_mut();
211        write.buf.extend_from_slice(data.as_ref());
212        write.waker.wake();
213    }
214
215    /// Sets how many bytes the peer may write before becoming blocked.
216    ///
217    /// Increasing the capacity wakes the peer's pending writer.
218    pub fn remote_buffer_cap(&self, cap: usize) {
219        // change cap
220        self.local.lock().unwrap().borrow_mut().buf_cap = cap;
221        // wake remote
222        self.remote.lock().unwrap().borrow().waker.wake();
223    }
224
225    /// Takes all bytes currently readable by this endpoint without waiting.
226    pub fn read_any(&self) -> Bytes {
227        self.local.lock().unwrap().borrow_mut().buf.take()
228    }
229
230    /// Waits for readable data or peer closure, then takes all available bytes.
231    pub async fn read(&self) -> Result<Bytes, io::Error> {
232        if self.local.lock().unwrap().borrow().buf.is_empty() {
233            poll_fn(|cx| {
234                let guard = self.local.lock().unwrap();
235                let read = guard.borrow_mut();
236                if read.buf.is_empty() {
237                    let closed = match self.tp {
238                        Type::Client | Type::ClientClone => {
239                            self.is_server_dropped() || read.is_closed()
240                        }
241                        Type::Server | Type::ServerClone => self.is_client_dropped(),
242                    };
243                    if closed {
244                        Poll::Ready(())
245                    } else {
246                        *read.waker.0.lock().unwrap().borrow_mut() = Some(cx.waker().clone());
247                        drop(read);
248                        drop(guard);
249                        Poll::Pending
250                    }
251                } else {
252                    Poll::Ready(())
253                }
254            })
255            .await;
256        }
257        Ok(self.local.lock().unwrap().borrow_mut().buf.take())
258    }
259
260    /// Polls a transport read into `buf`.
261    ///
262    /// Returns the number of bytes copied, zero on simulated peer closure, or
263    /// `Pending` when no input is available.
264    ///
265    /// # Panics
266    ///
267    /// Panics if data is available but `buf` has no remaining capacity.
268    pub fn poll_read_buf(
269        &self,
270        cx: &mut Context<'_>,
271        buf: &mut BytesMut,
272    ) -> Poll<io::Result<usize>> {
273        let guard = self.local.lock().unwrap();
274        let mut ch = guard.borrow_mut();
275        *ch.waker.0.lock().unwrap().borrow_mut() = Some(cx.waker().clone());
276
277        if !ch.buf.is_empty() {
278            let size = std::cmp::min(ch.buf.len(), buf.remaining_mut());
279            assert!(size > 0, "Supplied buffer is zero sized");
280            let b = ch.buf.split_to(size);
281            buf.put_slice(&b);
282            return Poll::Ready(Ok(size));
283        }
284
285        match mem::take(&mut ch.read) {
286            IoTestState::Ok | IoTestState::Pending => Poll::Pending,
287            IoTestState::Close => {
288                ch.read = IoTestState::Close;
289                Poll::Ready(Ok(0))
290            }
291            IoTestState::Err(e) => Poll::Ready(Err(e)),
292        }
293    }
294
295    /// Polls a transport write from `buf`.
296    ///
297    /// The write is limited by the capacity configured with
298    /// [`remote_buffer_cap`](Self::remote_buffer_cap).
299    pub fn poll_write_buf(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
300        let guard = self.remote.lock().unwrap();
301        let mut ch = guard.borrow_mut();
302
303        match mem::take(&mut ch.write) {
304            IoTestState::Ok => {
305                let cap = cmp::min(buf.len(), ch.buf_cap);
306                if cap > 0 {
307                    ch.buf.extend(&buf[..cap]);
308                    ch.buf_cap -= cap;
309                    ch.flags.remove(IoTestFlags::FLUSHED);
310                    ch.waker.wake();
311                    Poll::Ready(Ok(cap))
312                } else {
313                    *self
314                        .local
315                        .lock()
316                        .unwrap()
317                        .borrow_mut()
318                        .waker
319                        .0
320                        .lock()
321                        .unwrap()
322                        .borrow_mut() = Some(cx.waker().clone());
323                    Poll::Pending
324                }
325            }
326            IoTestState::Close => Poll::Ready(Ok(0)),
327            IoTestState::Pending => {
328                *self
329                    .local
330                    .lock()
331                    .unwrap()
332                    .borrow_mut()
333                    .waker
334                    .0
335                    .lock()
336                    .unwrap()
337                    .borrow_mut() = Some(cx.waker().clone());
338                Poll::Pending
339            }
340            IoTestState::Err(e) => Poll::Ready(Err(e)),
341        }
342    }
343}
344
345impl Clone for IoTest {
346    fn clone(&self) -> Self {
347        let tp = match self.tp {
348            Type::Server => Type::ServerClone,
349            Type::Client => Type::ClientClone,
350            val => val,
351        };
352
353        IoTest {
354            tp,
355            local: self.local.clone(),
356            remote: self.remote.clone(),
357            state: self.state.clone(),
358            peer_addr: self.peer_addr,
359        }
360    }
361}
362
363impl Drop for IoTest {
364    fn drop(&mut self) {
365        let mut state = *self.state.lock().unwrap().borrow();
366        match self.tp {
367            Type::Server => state.server_dropped = true,
368            Type::Client => state.client_dropped = true,
369            _ => (),
370        }
371        *self.state.lock().unwrap().borrow_mut() = state;
372
373        let guard = self.remote.lock().unwrap();
374        let mut remote = guard.borrow_mut();
375        remote.read = IoTestState::Close;
376        remote.waker.wake();
377        log::debug!("drop remote socket");
378    }
379}
380
381impl IoStream for IoTest {
382    fn start(self, ctx: IoContext) -> Box<dyn Handle> {
383        let io = Rc::new(self);
384        ntex_util::spawn(run(io.clone(), ctx));
385        Box::new(io)
386    }
387}
388
389impl Handle for Rc<IoTest> {
390    fn query(&self, id: any::TypeId) -> Option<Box<dyn any::Any>> {
391        if id == any::TypeId::of::<types::PeerAddr>()
392            && let Some(addr) = self.peer_addr
393        {
394            return Some(Box::new(types::PeerAddr(addr)));
395        }
396        None
397    }
398}
399
400async fn run(io: Rc<IoTest>, ctx: IoContext) {
401    poll_fn(|cx| turn(&io, &ctx, cx)).await;
402
403    log::debug!("{}: Shuting down io", ctx.tag());
404
405    // shutdown WRITE side
406    io.local
407        .lock()
408        .unwrap()
409        .borrow_mut()
410        .flags
411        .insert(IoTestFlags::CLOSED);
412
413    log::debug!("{}: Shutdown complete", ctx.tag());
414    ctx.stopped(None);
415}
416
417fn turn(io: &IoTest, ctx: &IoContext, cx: &mut Context<'_>) -> Poll<()> {
418    let read = match ctx.poll_read_ready(cx) {
419        Poll::Ready(Readiness::Ready) => read(io, ctx, cx),
420        Poll::Ready(Readiness::Close | Readiness::Terminate) => Poll::Ready(()),
421        Poll::Pending => Poll::Pending,
422    };
423
424    let write = match ctx.poll_write_ready(cx) {
425        Poll::Ready(Readiness::Ready) => write(io, ctx, cx),
426        Poll::Ready(Readiness::Close | Readiness::Terminate) => Poll::Ready(()),
427        Poll::Pending => Poll::Pending,
428    };
429
430    if read.is_pending() && write.is_pending() {
431        Poll::Pending
432    } else {
433        Poll::Ready(())
434    }
435}
436
437fn write(io: &IoTest, ctx: &IoContext, cx: &mut Context<'_>) -> Poll<()> {
438    let result = ctx.with_write_dst(|buf| write_io(io, buf, cx, ctx));
439    if ctx.update_write_status(result) == IoTaskStatus::Stop {
440        Poll::Ready(())
441    } else {
442        Poll::Pending
443    }
444}
445
446fn read(io: &IoTest, ctx: &IoContext, cx: &mut Context<'_>) -> Poll<()> {
447    loop {
448        let mut pending = false;
449        let result = ctx.with_read_buf(|buf| {
450            let result = io.poll_read_buf(cx, buf);
451            pending = result.is_pending();
452            result
453        });
454        return match result {
455            IoTaskStatus::Io => {
456                if pending {
457                    Poll::Pending
458                } else {
459                    continue;
460                }
461            }
462            IoTaskStatus::Stop => Poll::Ready(()),
463            IoTaskStatus::Pause => Poll::Pending,
464        };
465    }
466}
467
468/// Flush write buffer to underlying I/O stream.
469pub(super) fn write_io(
470    io: &IoTest,
471    buf: &mut BytePages,
472    cx: &mut Context<'_>,
473    ctx: &IoContext,
474) -> io::Result<usize> {
475    let tag = ctx.tag();
476    let mut written = 0;
477
478    while let Some(mut page) = buf.take() {
479        log::debug!("{tag}: flushing framed transport: {}", page.len());
480
481        let result = io.poll_write_buf(cx, &page)?;
482        match result {
483            Poll::Ready(0) => {
484                log::trace!("{tag}: disconnected during flush, written {written}");
485                buf.prepend(page);
486                return Err(io::Error::new(
487                    io::ErrorKind::WriteZero,
488                    "failed to write frame to transport",
489                ));
490            }
491            Poll::Ready(n) => {
492                written += n;
493                page.advance_to(n);
494                buf.prepend(page);
495            }
496            Poll::Pending => {
497                buf.prepend(page);
498                break;
499            }
500        }
501    }
502
503    log::debug!("{tag}: flushed {written} bytes, remaining: {}", buf.len());
504    Ok(written)
505}
506
507#[cfg(test)]
508#[allow(clippy::redundant_clone)]
509mod tests {
510    use super::*;
511    use ntex_util::future::lazy;
512
513    #[ntex::test]
514    async fn basic() {
515        let (client, server) = IoTest::create();
516        assert_eq!(client.tp, Type::Client);
517        assert_eq!(client.clone().tp, Type::ClientClone);
518        assert_eq!(server.tp, Type::Server);
519        assert_eq!(server.clone().tp, Type::ServerClone);
520        assert!(format!("{server:?}").contains("IoTest"));
521        assert!(format!("{:?}", AtomicWaker::default()).contains("AtomicWaker"));
522
523        server.read_pending();
524        let mut buf = BytesMut::new();
525        let res = lazy(|cx| client.poll_read_buf(cx, &mut buf)).await;
526        assert!(res.is_pending());
527
528        server.read_pending();
529        let res = lazy(|cx| server.poll_write_buf(cx, b"123")).await;
530        assert!(res.is_pending());
531
532        assert!(!server.is_client_dropped());
533        drop(client);
534        assert!(server.is_client_dropped());
535
536        let server2 = server.clone();
537        assert!(!server2.is_server_dropped());
538        drop(server);
539        assert!(server2.is_server_dropped());
540
541        let res = lazy(|cx| server2.poll_write_buf(cx, b"123")).await;
542        assert!(res.is_pending());
543
544        let (client, _) = IoTest::create();
545        let addr: net::SocketAddr = "127.0.0.1:8080".parse().unwrap();
546        let client = crate::Io::from(client.set_peer_addr(addr));
547        let item = client.query::<crate::types::PeerAddr>();
548        assert!(format!("{item:?}").contains("QueryItem(127.0.0.1:8080)"));
549    }
550}