1#![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#[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 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 pub fn is_client_dropped(&self) -> bool {
128 self.state.lock().unwrap().borrow().client_dropped
129 }
130
131 pub fn is_server_dropped(&self) -> bool {
133 self.state.lock().unwrap().borrow().server_dropped
134 }
135
136 pub fn is_closed(&self) -> bool {
138 self.remote.lock().unwrap().borrow().is_closed()
139 }
140
141 #[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 pub fn read_pending(&self) {
150 self.remote.lock().unwrap().borrow_mut().read = IoTestState::Pending;
151 }
152
153 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 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 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 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 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 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 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 pub fn remote_buffer_cap(&self, cap: usize) {
219 self.local.lock().unwrap().borrow_mut().buf_cap = cap;
221 self.remote.lock().unwrap().borrow().waker.wake();
223 }
224
225 pub fn read_any(&self) -> Bytes {
227 self.local.lock().unwrap().borrow_mut().buf.take()
228 }
229
230 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 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 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 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
468pub(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}